Compare commits

...

1 Commits

Author SHA1 Message Date
Soumil Rathi ebc82245a5 fix(oss): cap v3 add ingestion fanout 2026-06-16 17:05:20 -07:00
4 changed files with 309 additions and 21 deletions
+75 -9
View File
@@ -84,6 +84,58 @@ const ENTITY_PARAMS = [
"runId",
];
const ADD_CONTEXT_SEARCH_MAX_CHARS = 15000;
const ADD_MAX_EXTRACTED_MEMORIES = 20;
const ADD_MEMORY_EMBEDDING_MAX_CHARS = 15000;
const ADD_MAX_ENTITY_LINKS_PER_EVENT = 100;
const SEARCH_QUERY_MAX_CHARS = 15000;
function tailCapText(text: string, maxChars: number): string {
if (text.length <= maxChars) return text;
return text.slice(-maxChars);
}
function capExtractedMemories<T>(memories: T[]): T[] {
if (memories.length <= ADD_MAX_EXTRACTED_MEMORIES) return memories;
console.warn(
`Extracted memory cap applied: extracted=${memories.length} processing=${ADD_MAX_EXTRACTED_MEMORIES}`,
);
return memories.slice(-ADD_MAX_EXTRACTED_MEMORIES);
}
function capMemoryEmbeddingTexts(texts: string[]): string[] {
const capped = texts.map((text) =>
tailCapText(text, ADD_MEMORY_EMBEDDING_MAX_CHARS),
);
const cappedCount = texts.filter(
(text, index) => text !== capped[index],
).length;
if (cappedCount > 0) {
console.warn(
`Memory embedding input cap applied: capped_count=${cappedCount} max_chars=${ADD_MEMORY_EMBEDDING_MAX_CHARS}`,
);
}
return capped;
}
function capSearchQuery(query: string): string {
const capped = tailCapText(query, SEARCH_QUERY_MAX_CHARS);
if (capped.length !== query.length) {
console.warn(
`Search query cap applied: original_chars=${query.length} searched_chars=${capped.length}`,
);
}
return capped;
}
function capEntityKeys(keys: string[]): string[] {
if (keys.length <= ADD_MAX_ENTITY_LINKS_PER_EVENT) return keys;
console.warn(
`ADD entity-linking cap applied: extracted=${keys.length} processing=${ADD_MAX_ENTITY_LINKS_PER_EVENT}`,
);
return keys.slice(0, ADD_MAX_ENTITY_LINKS_PER_EVENT);
}
/**
* Validates that no top-level entity parameters are passed in config.
* @throws Error if entity params are found at top level
@@ -734,7 +786,16 @@ export class Memory {
const parsedMessages = messages.map((m) => m.content).join("\n");
// Phase 1: Existing memory retrieval
const queryEmbedding = await this.embedder.embed(parsedMessages);
const searchMessages = tailCapText(
parsedMessages,
ADD_CONTEXT_SEARCH_MAX_CHARS,
);
if (searchMessages.length !== parsedMessages.length) {
console.warn(
`ADD context-search embedding capped: original_chars=${parsedMessages.length} embedded_chars=${searchMessages.length}`,
);
}
const queryEmbedding = await this.embedder.embed(searchMessages);
const existingResults = await this.vectorStore.search(
queryEmbedding,
10,
@@ -805,6 +866,7 @@ export class Memory {
console.error("Error parsing extraction response:", e);
extractedMemories = [];
}
extractedMemories = capExtractedMemories(extractedMemories);
if (extractedMemories.length === 0) {
// Save messages even if nothing extracted
@@ -826,17 +888,21 @@ export class Memory {
const memTexts = extractedMemories
.map((m) => m.text ?? "")
.filter((t) => t.length > 0);
const memEmbeddingTexts = capMemoryEmbeddingTexts(memTexts);
let embedMap: Record<string, number[]> = {};
try {
const memEmbeddingsList = await this.embedder.embedBatch(memTexts);
const memEmbeddingsList =
await this.embedder.embedBatch(memEmbeddingTexts);
for (let i = 0; i < memTexts.length; i++) {
embedMap[memTexts[i]] = memEmbeddingsList[i];
}
} catch {
// Fallback: embed individually
for (const text of memTexts) {
for (let i = 0; i < memTexts.length; i++) {
const text = memTexts[i];
const embeddingText = memEmbeddingTexts[i];
try {
embedMap[text] = await this.embedder.embed(text);
embedMap[text] = await this.embedder.embed(embeddingText);
} catch (e) {
console.warn(`Failed to embed memory text: ${e}`);
}
@@ -1007,9 +1073,8 @@ export class Memory {
const orderedKeys = Object.keys(globalEntities);
if (orderedKeys.length > 0) {
const entityTexts = orderedKeys.map(
(k) => globalEntities[k].entityText,
);
const entityKeys = capEntityKeys(orderedKeys);
const entityTexts = entityKeys.map((k) => globalEntities[k].entityText);
// 7b: Single batch embed for all unique entities
let entityEmbeddings: (number[] | null)[];
@@ -1029,9 +1094,9 @@ export class Memory {
// Filter out entities with failed embeddings
const valid: Array<{ index: number; key: string }> = [];
for (let i = 0; i < orderedKeys.length; i++) {
for (let i = 0; i < entityKeys.length; i++) {
if (entityEmbeddings[i] !== null) {
valid.push({ index: i, key: orderedKeys[i] });
valid.push({ index: i, key: entityKeys[i] });
}
}
@@ -1256,6 +1321,7 @@ export class Memory {
}
const searchStartMs = Date.now();
query = capSearchQuery(query);
// Step 1: Preprocess query
const queryLemmatized = lemmatizeForBm25(query);
+76
View File
@@ -176,3 +176,79 @@ describe("Memory - add()", () => {
);
});
});
describe("Memory - ingestion caps", () => {
test("caps add context-search embedding to the recent tail", async () => {
const memory = createMemory({
vectorStore: {
provider: "memory",
config: {
collectionName: `test-add-context-cap-${Date.now()}`,
dimension: 1536,
dbPath: ":memory:",
},
},
});
const embedSpy = jest.spyOn((memory as any).embedder, "embed");
const longText = `old-${"x".repeat(15000)}`;
await memory.add(longText, { userId: `cap_user_${Date.now()}` });
expect(embedSpy.mock.calls[0][0]).toBe(longText.slice(-15000));
await memory.reset();
});
test("caps extracted memories to the latest 20 before embedding", async () => {
const memory = createMemory({
vectorStore: {
provider: "memory",
config: {
collectionName: `test-add-memory-cap-${Date.now()}`,
dimension: 1536,
dbPath: ":memory:",
},
},
});
(memory as any).llm.generateResponse = jest.fn().mockResolvedValue(
JSON.stringify({
memory: Array.from({ length: 25 }, (_, i) => ({
id: String(i),
text: `memory-${i}`,
})),
}),
);
const embedBatchSpy = jest.spyOn((memory as any).embedder, "embedBatch");
const result: SearchResult = await memory.add("remember these", {
userId: `cap_user_${Date.now()}`,
});
expect(result.results).toHaveLength(20);
expect(embedBatchSpy.mock.calls[0][0]).toEqual(
Array.from({ length: 20 }, (_, i) => `memory-${i + 5}`),
);
await memory.reset();
});
test("caps search query before embedding", async () => {
const memory = createMemory({
vectorStore: {
provider: "memory",
config: {
collectionName: `test-search-cap-${Date.now()}`,
dimension: 1536,
dbPath: ":memory:",
},
},
});
const embedSpy = jest.spyOn((memory as any).embedder, "embed");
const longQuery = `old-${"q".repeat(15000)}`;
await memory.search(longQuery, {
filters: { user_id: `cap_user_${Date.now()}` },
});
expect(embedSpy.mock.calls[0][0]).toBe(longQuery.slice(-15000));
await memory.reset();
});
});
+89 -12
View File
@@ -80,6 +80,63 @@ warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*swigva
# Initialize logger early for util functions
logger = logging.getLogger(__name__)
ADD_CONTEXT_SEARCH_MAX_CHARS = 15000
ADD_MAX_EXTRACTED_MEMORIES = 20
ADD_MEMORY_EMBEDDING_MAX_CHARS = 15000
ADD_MAX_ENTITY_LINKS_PER_EVENT = 100
SEARCH_QUERY_MAX_CHARS = 15000
def _tail_cap_text(text: str, max_chars: int) -> str:
if len(text) <= max_chars:
return text
return text[-max_chars:]
def _cap_extracted_memories(extracted_memories: list) -> list:
if len(extracted_memories) <= ADD_MAX_EXTRACTED_MEMORIES:
return extracted_memories
logger.warning(
"Extracted memory cap applied: extracted=%d processing=%d",
len(extracted_memories),
ADD_MAX_EXTRACTED_MEMORIES,
)
return extracted_memories[-ADD_MAX_EXTRACTED_MEMORIES:]
def _cap_memory_embedding_texts(mem_texts: list[str]) -> list[str]:
capped = [_tail_cap_text(t, ADD_MEMORY_EMBEDDING_MAX_CHARS) for t in mem_texts]
capped_count = sum(1 for original, capped_text in zip(mem_texts, capped) if len(original) != len(capped_text))
if capped_count:
logger.warning(
"Memory embedding input cap applied: capped_count=%d max_chars=%d",
capped_count,
ADD_MEMORY_EMBEDDING_MAX_CHARS,
)
return capped
def _cap_search_query(query: str) -> str:
capped = _tail_cap_text(query, SEARCH_QUERY_MAX_CHARS)
if len(capped) != len(query):
logger.warning(
"Search query cap applied: original_chars=%d searched_chars=%d",
len(query),
len(capped),
)
return capped
def _cap_entity_keys(entity_keys: list[str]) -> list[str]:
if len(entity_keys) <= ADD_MAX_ENTITY_LINKS_PER_EVENT:
return entity_keys
logger.warning(
"ADD entity-linking cap applied: extracted=%d processing=%d",
len(entity_keys),
ADD_MAX_ENTITY_LINKS_PER_EVENT,
)
return entity_keys[:ADD_MAX_ENTITY_LINKS_PER_EVENT]
# Fields that hold runtime auth/connection objects and must be preserved.
# These are non-serializable objects (e.g. AWSV4SignerAuth, RequestsHttpConnection)
@@ -804,9 +861,16 @@ class Memory(MemoryBase):
# Phase 1: Existing memory retrieval
search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v}
query_embedding = self.embedding_model.embed(parsed_messages, "search")
search_messages = _tail_cap_text(parsed_messages, ADD_CONTEXT_SEARCH_MAX_CHARS)
if len(search_messages) != len(parsed_messages):
logger.warning(
"ADD context-search embedding capped: original_chars=%d embedded_chars=%d",
len(parsed_messages),
len(search_messages),
)
query_embedding = self.embedding_model.embed(search_messages, "search")
existing_results = self.vector_store.search(
query=parsed_messages,
query=search_messages,
vectors=query_embedding,
top_k=10,
filters=search_filters,
@@ -860,6 +924,7 @@ class Memory(MemoryBase):
except Exception as e:
logger.error(f"Error parsing extraction response: {e}")
extracted_memories = []
extracted_memories = _cap_extracted_memories(extracted_memories)
if not extracted_memories:
# Save messages even if nothing extracted
@@ -868,15 +933,16 @@ class Memory(MemoryBase):
# Phase 3: Batch embed all extracted memory texts
mem_texts = [m.get("text", "") for m in extracted_memories if m.get("text")]
mem_embedding_texts = _cap_memory_embedding_texts(mem_texts)
try:
mem_embeddings_list = self.embedding_model.embed_batch(mem_texts, "add")
mem_embeddings_list = self.embedding_model.embed_batch(mem_embedding_texts, "add")
embed_map = dict(zip(mem_texts, mem_embeddings_list))
except Exception:
# Fallback: embed individually
embed_map = {}
for text in mem_texts:
for text, embedding_text in zip(mem_texts, mem_embedding_texts):
try:
embed_map[text] = self.embedding_model.embed(text, "add")
embed_map[text] = self.embedding_model.embed(embedding_text, "add")
except Exception as e:
logger.warning(f"Failed to embed memory text: {e}")
@@ -978,7 +1044,7 @@ class Memory(MemoryBase):
global_entities[key] = [entity_type, entity_text, {memory_id}]
if global_entities:
ordered_keys = list(global_entities.keys())
ordered_keys = _cap_entity_keys(list(global_entities.keys()))
entity_texts = [global_entities[k][1] for k in ordered_keys]
# 7b: Single batch embed for all unique entities
@@ -1478,6 +1544,7 @@ class Memory(MemoryBase):
# Guard against None threshold (backward compat)
if threshold is None:
threshold = 0.1
query = _cap_search_query(query)
# Step 1: Preprocess query
query_lemmatized = lemmatize_for_bm25(query)
@@ -2310,10 +2377,17 @@ class AsyncMemory(MemoryBase):
# Phase 1: Existing memory retrieval
search_filters = {k: v for k, v in effective_filters.items() if k in ("user_id", "agent_id", "run_id") and v}
query_embedding = await asyncio.to_thread(self.embedding_model.embed, parsed_messages, "search")
search_messages = _tail_cap_text(parsed_messages, ADD_CONTEXT_SEARCH_MAX_CHARS)
if len(search_messages) != len(parsed_messages):
logger.warning(
"ADD context-search embedding capped: original_chars=%d embedded_chars=%d",
len(parsed_messages),
len(search_messages),
)
query_embedding = await asyncio.to_thread(self.embedding_model.embed, search_messages, "search")
existing_results = await asyncio.to_thread(
self.vector_store.search,
query=parsed_messages,
query=search_messages,
vectors=query_embedding,
top_k=10,
filters=search_filters,
@@ -2368,6 +2442,7 @@ class AsyncMemory(MemoryBase):
except Exception as e:
logger.error(f"Error parsing extraction response (async): {e}")
extracted_memories = []
extracted_memories = _cap_extracted_memories(extracted_memories)
if not extracted_memories:
await asyncio.to_thread(self.db.save_messages, messages, session_scope)
@@ -2375,14 +2450,15 @@ class AsyncMemory(MemoryBase):
# Phase 3: Batch embed all extracted memory texts
mem_texts = [m.get("text", "") for m in extracted_memories if m.get("text")]
mem_embedding_texts = _cap_memory_embedding_texts(mem_texts)
try:
mem_embeddings_list = await asyncio.to_thread(self.embedding_model.embed_batch, mem_texts, "add")
mem_embeddings_list = await asyncio.to_thread(self.embedding_model.embed_batch, mem_embedding_texts, "add")
embed_map = dict(zip(mem_texts, mem_embeddings_list))
except Exception:
embed_map = {}
for text in mem_texts:
for text, embedding_text in zip(mem_texts, mem_embedding_texts):
try:
embed_map[text] = await asyncio.to_thread(self.embedding_model.embed, text, "add")
embed_map[text] = await asyncio.to_thread(self.embedding_model.embed, embedding_text, "add")
except Exception as e:
logger.warning(f"Failed to embed memory text (async): {e}")
@@ -2485,7 +2561,7 @@ class AsyncMemory(MemoryBase):
global_entities[key] = [entity_type, entity_text, {memory_id}]
if global_entities:
ordered_keys = list(global_entities.keys())
ordered_keys = _cap_entity_keys(list(global_entities.keys()))
entity_texts = [global_entities[k][1] for k in ordered_keys]
# 7b: Batch embed entities
@@ -2990,6 +3066,7 @@ class AsyncMemory(MemoryBase):
async def _search_vector_store(self, query, filters, limit, threshold=0.1, explain=False):
if threshold is None:
threshold = 0.1
query = _cap_search_query(query)
# Step 1: Preprocess query (CPU-bound)
query_lemmatized = await asyncio.to_thread(lemmatize_for_bm25, query)
+69
View File
@@ -0,0 +1,69 @@
import logging
from unittest.mock import Mock
from mem0.memory.main import (
ADD_CONTEXT_SEARCH_MAX_CHARS,
ADD_MAX_ENTITY_LINKS_PER_EVENT,
ADD_MAX_EXTRACTED_MEMORIES,
ADD_MEMORY_EMBEDDING_MAX_CHARS,
Memory,
SEARCH_QUERY_MAX_CHARS,
_cap_entity_keys,
_cap_extracted_memories,
_cap_memory_embedding_texts,
_tail_cap_text,
)
def test_tail_cap_text_keeps_recent_suffix():
text = "old-" + ("x" * ADD_CONTEXT_SEARCH_MAX_CHARS)
capped = _tail_cap_text(text, ADD_CONTEXT_SEARCH_MAX_CHARS)
assert capped == text[-ADD_CONTEXT_SEARCH_MAX_CHARS:]
def test_cap_extracted_memories_keeps_last_memories(caplog):
memories = [{"text": f"memory-{i}"} for i in range(ADD_MAX_EXTRACTED_MEMORIES + 3)]
with caplog.at_level(logging.WARNING):
capped = _cap_extracted_memories(memories)
assert capped == memories[-ADD_MAX_EXTRACTED_MEMORIES:]
assert any("Extracted memory cap applied" in record.message for record in caplog.records)
def test_cap_memory_embedding_texts_keeps_original_text_untouched(caplog):
original = "old-" + ("z" * ADD_MEMORY_EMBEDDING_MAX_CHARS)
with caplog.at_level(logging.WARNING):
capped = _cap_memory_embedding_texts([original])
assert capped == [original[-ADD_MEMORY_EMBEDDING_MAX_CHARS:]]
assert original.startswith("old-")
assert any("Memory embedding input cap applied" in record.message for record in caplog.records)
def test_cap_entity_keys_keeps_first_100_entities(caplog):
keys = [f"entity-{i}" for i in range(ADD_MAX_ENTITY_LINKS_PER_EVENT + 3)]
with caplog.at_level(logging.WARNING):
capped = _cap_entity_keys(keys)
assert capped == keys[:ADD_MAX_ENTITY_LINKS_PER_EVENT]
assert any("ADD entity-linking cap applied" in record.message for record in caplog.records)
def test_search_vector_store_caps_query_before_embedding():
memory = object.__new__(Memory)
memory.embedding_model = Mock()
memory.embedding_model.embed.return_value = [0.1, 0.2, 0.3]
memory.vector_store = Mock()
memory.vector_store.search.return_value = []
memory.vector_store.keyword_search.return_value = None
memory._compute_entity_boosts = Mock(return_value={})
query = "old-" + ("q" * SEARCH_QUERY_MAX_CHARS)
Memory._search_vector_store(memory, query, filters={}, limit=5)
memory.embedding_model.embed.assert_called_once_with(query[-SEARCH_QUERY_MAX_CHARS:], "search")