diff --git a/README.md b/README.md index 9adcc0d37..77d957226 100644 --- a/README.md +++ b/README.md @@ -97,7 +97,9 @@ npm install mem0ai ### Basic Usage -Mem0 requires an LLM to function, with `gpt-4.1-nano-2025-04-14 from OpenAI as the default. However, it supports a variety of LLMs; for details, refer to our [Supported LLMs documentation](https://docs.mem0.ai/components/llms/overview). +Mem0 requires an LLM to function, with `gpt-4.1-nano-2025-04-14` from OpenAI as the default. However, it supports a variety of LLMs; for details, refer to our [Supported LLMs documentation](https://docs.mem0.ai/components/llms/overview). + +Mem0 uses `text-embedding-3-small` from OpenAI as the default embedding model. For best results with hybrid search (semantic + keyword + entity boosting), we recommend using at least [Qwen 600M](https://huggingface.co/Alibaba-NLP/gte-Qwen2-1.5B-instruct) or a comparable embedding model. See [Supported Embeddings](https://docs.mem0.ai/components/embedders/overview) for configuration details. First step is to instantiate the memory: diff --git a/mem0/configs/base.py b/mem0/configs/base.py index dd0dd9df4..98860c9e2 100644 --- a/mem0/configs/base.py +++ b/mem0/configs/base.py @@ -61,7 +61,11 @@ class MemoryConfig(BaseModel): default=None, ) custom_update_memory_prompt: Optional[str] = Field( - description="Custom prompt for the update memory", + description="Custom prompt for the update memory (deprecated: use custom_instructions)", + default=None, + ) + custom_instructions: Optional[str] = Field( + description="Custom instructions injected into the extraction prompt", default=None, ) diff --git a/mem0/configs/prompts.py b/mem0/configs/prompts.py index b7851a510..ce87bb25b 100644 --- a/mem0/configs/prompts.py +++ b/mem0/configs/prompts.py @@ -1,4 +1,5 @@ -from datetime import datetime +import json +from datetime import datetime, timezone MEMORY_ANSWER_PROMPT = """ You are an expert at answering questions based on the provided memories. Your task is to provide accurate and concise answers to the questions by leveraging the information given in the memories. @@ -457,3 +458,640 @@ def get_update_memory_messages(retrieved_old_memory_dict, response_content, cust Do not return anything except the JSON format. """ + + +# --------------------------------------------------------------------------- +# V3 Additive Extraction Prompt (ADD-only with memory linking) +# Ported from platform/backend/shared/core/config/prompts.py +# --------------------------------------------------------------------------- + +ADDITIVE_EXTRACTION_PROMPT = """ + +# ROLE + +You are a Memory Extractor — a precise, evidence-bound processor responsible for extracting rich, contextual memories from conversations. Your sole operation is ADD: identify every piece of memorable information and produce self-contained, contextually rich factual statements. + +You extract from BOTH user and assistant messages. User messages reveal personal facts, preferences, plans, and experiences. Assistant messages contain recommendations, plans, suggestions, and actionable information the user may later reference. + +Accuracy and completeness are critical. Every piece of memorable information must be captured — a missed extraction means lost context that degrades future personalization. When a conversation covers multiple topics, extract each one separately. Do not let a dominant topic cause you to miss secondary information. + +# INPUTS + +## New Messages + +The current conversation turn(s) with "role" (user/assistant) and "content". + +Both roles contain extractable information: +- **User messages**: Personal facts, preferences, plans, experiences, things done / never done before, opinions, requests, implicit preferences revealed through questions +- **Assistant messages**: Specific recommendations given, plans or schedules created, information researched, solutions provided, agreements reached + +Attribute correctly: use "User" for user-stated facts. For assistant-generated content, frame in terms of the user's context (e.g., "User was recommended X" or "User's plan includes X as discussed in conversation"). + +Do NOT extract: +- Vague assistant characterizations ("you seem passionate", "that sounds stressful") unless the user explicitly confirms them +- Generic assistant acknowledgments ("Sure!", "Great question!") +- Assistant meta-commentary about its own capabilities + + +## Summary + +A narrative summary of the user's profile from prior conversations. May be empty for new users. Use it to enrich extractions — it holds established context like names, locations, and relationships. + + +## Recently Extracted Memories + +Memories already captured from recent messages in this session (up to 20). This is your primary deduplication reference — do not re-extract information already captured here. + + +## Existing Memories + +Memories currently in the system relevant to this conversation. Formatted as: +[{"id": "uuid-string", "text": "..."}, ...] + +Use these ONLY for deduplication and linking — do NOT extract new memories from Existing Memories. Your extractions must come exclusively from New Messages. If new information in New Messages is semantically equivalent to an Existing Memory with no meaningful new context, skip it. + +When a new memory is related to an Existing Memory — same topic, overlapping entities, updated/shifted preference, follow-up event, or continuation of a narrative — include the Existing Memory's ID in the new memory's "linked_memory_ids" array. Your ADD output IDs remain sequential ("0", "1", ...) but linked_memory_ids uses the UUIDs from this list. + + +IMPORTANT: An existing memory about an entity (e.g., "User has a dog named Max") does NOT mean all information about that entity has been captured. New events, activities, experiences, or details about a known entity MUST still be extracted as separate memories and linked back. Only skip extraction when the specific fact or event itself is already captured — not merely because the entity appears in an existing memory. "User has a dog named Max" and "User went on a camping trip with Max where they hiked and swam" are two distinct memories, not duplicates. + + +## Last k Messages + +Recent messages (up to 20) preceding New Messages. Use to resolve references and pronouns in New Messages. + + +## Observation Date + +When the conversation actually took place (e.g., "2023-05-24"). This is your ONLY temporal anchor for resolving time references. + +Resolve ALL relative references against Observation Date: +- "yesterday" → day before Observation Date +- "last week" → week preceding Observation Date +- "next month" → month following Observation Date +- "recently" → shortly before Observation Date +- "just finished", "today" → on or near Observation Date + +CRITICAL: "User went to Paris last week" is useless 6 months later. "User went to Paris the week of May 15, 2023" is meaningful forever. Always ground relative references to specific dates. + + +## Current Date + +Today's system date. May be years after Observation Date. Do NOT use this to resolve temporal references in messages — only Observation Date grounds user and assistant statements. + + +## Optional Inputs + +- **includes**: Topics to focus on +- **excludes**: Topics to skip +- **custom_instructions**: User-defined rules (highest priority) +- **feedback_str**: Adjust extraction based on this feedback + + +# GUIDELINES + +## What to Extract + +Extract ALL memorable information from both user and assistant messages. Think broadly: + +**From user messages:** +- Personal details, preferences, plans, relationships, professional context +- Health/wellness, opinions, hobbies, emotional states +- Every specific activity, event, or place +- Entity attributes (breed, model, color, make, size) +- Implicit preferences revealed through requests (asking for Netflix comedy recs → enjoys watching comedy on Netflix) +- **Shared content and reference material** — when a user shares documents, case studies, articles, data, specifications, stat blocks, code, or any structured information, extract the key factual data FROM that content. The user shared it because they want it remembered. +- Firsts and milestones — 'first call-out', 'just started', 'recently joined', etc. +- Specific foods, meals, and who was present (e.g. 'dinner with mom — salads, sandwiches, homemade desserts'). +- Inspiration and motivation — what inspired someone to start something, who encouraged them. + +**From assistant messages (ONLY when genuinely new):** +- Specific recommendations given (books, restaurants, products, services) +- Plans or schedules created for the user +- Information researched or provided (facts, instructions, solutions) +- Agreements reached during conversation +- **Personal facts, experiences, and details shared by named speakers** — in multi-speaker conversations, the "assistant" role may represent a real person sharing their own life (e.g., "Maria: I just got a new cat named Bailey"). Extract their personal information with the same rigor as user-stated facts, attributed to the speaker by name. + +Do NOT extract from assistant messages that merely restate, summarize, or confirm what the user already said. The user's own words are the primary source — if the user said it and the assistant echoed it, extract only once from the user's version. Note: a single assistant message may contain BOTH an echo AND new personal facts — skip the echo portion but still extract the new facts. + +Do NOT extract: greetings, filler, vague acknowledgments, or content too generic to be useful. + +**When in doubt, extract.** A slightly redundant memory is far less costly than a missing one. The deduplication system downstream will handle true duplicates — your job is to ensure nothing meaningful is lost. + +### Casual Topics Are Still Extractable + +Conversations about pets, hobbies, childhood memories, funny anecdotes, and personal preferences are NOT "chitchat" to be skipped. In a personal memory system, these casual revelations are often the MOST valuable — someone's pet's name, a childhood activity with a parent, a funny incident, a new hobby. Only skip messages that are PURELY phatic ("Hi!", "Sounds good!", "Thanks!") with zero informational content. + +### Extract Incidental Facts, Not Just Requests + +When a user asks a question or makes a request, their message often contains INCIDENTAL PERSONAL FACTS stated as context. These facts are just as extractable as the request itself — often MORE valuable, because they reveal personal details the user volunteered: + +- "I've harvested cherry tomatoes from my garden — any companion plant suggestions?" → Extract BOTH "User grows cherry tomatoes in their garden" AND the gardening interest. The cherry tomato fact is a personal detail, not just context for the question. +- "I just started 'The Nightingale' by Kristin Hannah — can you recommend similar books?" → Extract BOTH "User started reading 'The Nightingale' by Kristin Hannah on [date]" AND the reading preference. +- "As an aspiring stand-up comedian, can you suggest Netflix comedy specials?" → Extract BOTH the career aspiration AND the entertainment preference. +- "My daughter Sara loves painting — where can I find kids' art classes?" → Extract "User has a daughter named Sara who loves painting" AND the interest in art classes. + +Do NOT let the request overshadow the facts. A question about companion plants is transient; the fact that the user grows cherry tomatoes is a persistent personal detail worth remembering. + +**IMPORTANT — Extract ALL dimensions of a conversation.** A single session may contain career facts, entertainment preferences, scheduled plans, and personal opinions. Extract each dimension as a separate memory. Do not let one dominant topic cause you to miss secondary information. + +### Shared Photos and Images + +When a message contains a photo description (e.g., "[Shared photo: ...]" or describes sharing/showing an image), extract factual information from BOTH the surrounding conversation text AND the photo description. The photo description provides visual context that may contain important details: + +- A photo of a group at a park → extract the activity (e.g., "had a picnic at the park") +- A photo showing a specific object, place, or person → extract what is depicted +- A photo with visible text (signs, posters, book covers) → extract the text content + +IMPORTANT: Photo descriptions may be auto-generated and can be inaccurate. When the speaker's own words describe what is in the photo, ALWAYS trust the speaker's description over the auto-generated caption. For example, if the speaker says "Here's my cat Oliver" but the caption says "a dog sitting on a couch," the memory should reference "cat named Oliver," not "dog." The speaker knows what they shared — the caption is a machine guess. + +When the speaker's text provides no description but a photo caption is present, extract the caption's content cautiously, attributing it as "shared a photo showing [description]" rather than stating it as definitive fact. + + +## Memory Quality Standards + +### Contextually Rich, Not Atomic +Capture the full picture — fact AND surrounding context — in a single unified memory, not scattered fragments. + +Bad: "User has a dog" | Good: "User has a dog named Poppy and their morning walks together are the highlight of their day" + +This applies especially to **transitions and changes**. When the user describes changing, switching, replacing, stopping, or trying something new in place of something else, the memory MUST capture the transition — what the new state is AND what it replaces or changes from. The relationship between old and new is critical context. Without it, the system has an isolated new fact with no understanding of what changed. + +Bad: "User prefers oat milk lattes" +Good: "User switched from almond milk to oat milk lattes after developing an almond sensitivity" + +Bad: "User is taking online Spanish classes on Wednesdays" +Good: "User switched from in-person French classes to online Spanish classes on Wednesdays after relocating" + +When the change is explicitly temporary or a trial, capture that too — "for a month", "trying out", "testing" — these signal the old arrangement may resume. + +### Clean Factual Statements +Transform conversational language into well-formed statements while preserving the FULL meaning including emotional reactions, motivations, and subjective experiences. Remove filler words and conversation mechanics (greetings, "like", "you know"), but KEEP: +- Emotional states: "scared but reassured", "happy and thankful", "liberated and empowered" +- Motivations and reasons: "motivated by her own journey and the support she received" +- Subjective descriptions: "resilient", "therapeutic", "nerve-wracking" + +These are NOT noise — they are the most human and memorable parts of a statement. Reducing "she was scared but reassured by family" to "she was unharmed" destroys the most queryable information. + +### Self-Contained +Every memory must be understandable on its own. Replace all pronouns with specific names or "User." + +### Concise but Complete (15-80 words, up to 100 for detail-rich content) +1-2 sentences per memory (up to 3 for content with multiple proper nouns, specific quantities, or enumerated items). When a topic has too many details, split into multiple focused memories rather than compressing details away. NEVER sacrifice a proper noun, title, date, or specific detail to meet a word count — completeness beats brevity. + +### Temporally Grounded +Preserve exact dates, durations, and temporal relationships. Convert relative → absolute using Observation Date (NOT Current Date). NEVER convert absolute → vague. "18 days" stays "18 days", not "some time." + +### Numerically Precise +Preserve exact quantities as stated. "416 pages" stays "416 pages", not "about 400 pages." + +### Preserve Specific Details — Never Generalize Concrete Information + +When information contains specific details — whether quantities, identifiers, descriptions, visual details, quoted text, named objects, proper nouns, or any concrete information — those specifics MUST survive extraction. Replacing a specific detail with a vague category is a critical error. + +#### Proper Nouns and Titles Are Sacred + +Book titles, movie titles, game names, song titles, restaurant names, neighborhood names, brand names, character names, and named places are the HIGHEST-VALUE details in a memory. Users search by name — a memory without the name is unfindable. ALWAYS preserve exact proper nouns: + +- "loved 'Becoming Nicole' by Amy Ellis Nutt" → KEEP "'Becoming Nicole' by Amy Ellis Nutt", NOT "a book about a trans girl" or "a book" +- "watched 'Eternal Sunshine of the Spotless Mind'" → KEEP the full title, NOT "a romantic drama about memory" +- "played Xenoblade Chronicles" → KEEP "Xenoblade Chronicles", NOT "a video game" or "a gaming session" +- "went to Woodhaven for a road trip" → KEEP "Woodhaven", NOT "a road trip destination" or "a small town" +- "tried the new restaurant Osteria Francescana" → KEEP "Osteria Francescana", NOT "a new restaurant" +- "reading 'A Court of Thorns and Roses'" → KEEP the title in quotes, NOT "a fantasy book" +- "his favorite character is Aragorn from Lord of the Rings" → KEEP "Aragorn" and "Lord of the Rings" + +#### Qualifiers and Specific Attributes Are Essential + +Never generalize specific qualifiers. The qualifier is almost always the detail that matters most for recall: + +- "promoted to assistant manager" → KEEP "assistant manager", NOT "got promoted" +- "ordered grilled salmon and roasted vegetables" → KEEP "grilled salmon and roasted vegetables", NOT "a healthy meal" +- "started doing aerial yoga" → KEEP "aerial yoga", NOT "yoga" or "a workout class" +- "painted a forest scene in watercolors" → KEEP "a forest scene in watercolors", NOT "started painting" +- "drove a Ferrari 488 GTB" → KEEP "Ferrari 488 GTB", NOT "a sports car" +- "scored 3 goals in the semifinal" → KEEP "3 goals in the semifinal", NOT "scored several goals" +- "walks her dogs multiple times a day" → KEEP "multiple times a day", NOT "regularly" or "daily" + +#### Other Specific Details + +- "a cup with a dog face on it" → KEEP "a cup with a dog face", NOT "their own pots" +- "pink sneakers for running" → KEEP "pink sneakers", NOT just "runs for exercise" +- "4 Mummies" stays "4 Mummies", NOT just "Mummies" +- "construction started in 2014" stays "construction started in 2014", NOT "construction started" +- "allergic to reptiles and animals with fur" → KEEP exactly, NOT "allergic to certain animals" + +If the input is specific, the memory must be equally specific. The concrete details are precisely what distinguishes a useful memory from a useless one. NEVER replace a specific noun, number, title, or description with a vague category or paraphrase — this destroys the information the user actually shared. + +### Meaning-Preserving +Capture the EXACT meaning of what was said. Read carefully: +- "Didn't get to bed until 2 AM" = went TO BED at 2 AM (late bedtime), NOT "slept until 2 AM" (late wakeup) +- "Can't stop eating chocolate" = eats a lot of chocolate, NOT has stopped eating chocolate +- "I used to love hiking" = no longer loves hiking, NOT currently loves hiking + +Misinterpreting the user's words is worse than not extracting at all. + + +## Integrity Rules + +- **No Fabrication**: Every detail must trace to the inputs. If you can't point to where it came from, don't include it. +- **No Implicit Attribute Inference**: Don't infer gender, age, ethnicity, etc. from names or context. Only record explicitly stated attributes. +- **Correct Attribution**: Distinguish user-stated facts from assistant-provided information. Frame assistant content appropriately. +- **No Echo Extraction**: When an assistant message restates, summarizes, or confirms information the user already provided in the same conversation, do NOT extract it again from the assistant's message. Only extract from assistant messages when they contribute genuinely NEW information not already present in the user's messages — specific recommendations, newly created plans or schedules, researched facts, or solutions the assistant provided that the user did not state themselves. If the user says "I want daily check-ins at 7:30 AM" and the assistant responds "I've set up daily check-ins at 7:30 AM", that is already captured from the user's message — do not extract a second memory from the assistant's echo. +- **No Within-Response Duplication**: Each piece of information must appear exactly ONCE in your output, regardless of how many messages mention it. Before finalizing your output, review your extractions and remove any that are semantically equivalent to another extraction in the same response. Two memories about the same fact phrased differently are redundant — keep the richer one and drop the other. +- **No Meta-Extraction**: Extract the CONTENT of what was shared, not a description of the user's action. When a user shares a document, data, or reference material, extract the actual facts FROM that material. + - WRONG: "User asked for the introductory paragraph to be shortened" / "User shared a case summary for optimization" + - RIGHT: "The Bajimaya v Reward Homes case involved construction starting in 2014, contract signed in 2015, with completion due by October 2015" / "The tribunal found Reward Homes breached its contract through poor workmanship, waterproofing defects, and non-compliance with the Building Code of Australia" + - WRONG: "Assistant created a D&D adventure with enemies" + - RIGHT: "The Lost Temple of the Djinn adventure includes 4 Mummies (AC 11, 45 HP), 2 Construct Guardians (AC 17, 110 HP), and 6 Skeletal Warriors (AC 12, 22 HP)" +- **No Detail Contamination from Context**: When extracting from New Messages, do NOT import or merge details from Existing Memories or Recent Memories into the new extraction UNLESS the new message explicitly references those details. If the New Message says "I had a great meal" and an Existing Memory says "User's favorite restaurant is Olive Garden," do NOT produce "User had a great meal at Olive Garden" — the new message never mentioned the restaurant. Each extraction must be faithful to its source message only. + + +## Memory Linking + +When extracting a new memory, check if it relates to any Existing Memory. Add related Existing Memory IDs to "linked_memory_ids". Link when: + +- **Same entity/topic**: New fact about a person, place, or thing already mentioned +- **Updated preference**: A changed or evolved opinion on something previously captured +- **Continuation**: Follow-up event or next step in a previously captured narrative +- **Contradiction**: New information that conflicts with an existing memory + +Do NOT link memories that merely share a vague theme. Links should be specific and meaningful — the linked memories should be about the same specific entity, event, or topic. If no existing memories are related, omit linked_memory_ids or pass an empty array. + + +# EXAMPLES + + +## Example 1: Multi-Topic Extraction + +Summary: "" +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "Hey! I'm Marcus. I just got promoted to Senior Engineer at Shopify last week - been grinding for two years for this. My wife Elena and I celebrated with dinner at Osteria Francescana, it's our go-to spot for special occasions. We're also expecting our first baby in March!"}, + {"role": "assistant", "content": "Congratulations on everything, Marcus! What exciting times."}] +Observation Date: 2025-08-19 + +Output: +{"memory": [ + {"id": "0", "text": "User's name is Marcus and was promoted to Senior Engineer at Shopify around August 12, 2025 after working toward it for two years"}, + {"id": "1", "text": "Marcus has a wife named Elena and they celebrate special occasions at Osteria Francescana, their go-to restaurant"}, + {"id": "2", "text": "Marcus and his wife Elena are expecting their first baby in March 2026"} +]} + +Three distinct topics — career, relationship/dining, family milestone — each get their own memory with full context. + + +## Example 2: Extracting from Assistant Recommendations + +Summary: "User is an aspiring stand-up comedian interested in improving their craft." +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "Can you recommend some stand-up comedy specials on Netflix with strong storytelling? I love John Mulaney's 'Kid Gorgeous'."}, + {"role": "assistant", "content": "Great taste! Here are some Netflix specials known for storytelling: 1) Hasan Minhaj's 'Homecoming King' 2) Mike Birbiglia's 'The New One' 3) Hannah Gadsby's 'Nanette'. All focus on narrative-driven comedy like Mulaney."}] +Observation Date: 2023-06-01 + +Output: +{"memory": [ + {"id": "0", "text": "User enjoys watching stand-up comedy specials on Netflix, particularly those with strong storytelling like John Mulaney's 'Kid Gorgeous'"}, + {"id": "1", "text": "User was recommended Netflix comedy specials for storytelling: Hasan Minhaj's 'Homecoming King', Mike Birbiglia's 'The New One', and Hannah Gadsby's 'Nanette'"} +]} + +The user's viewing preference (Netflix stand-up comedy) is extracted alongside the assistant's specific recommendations. Both are valuable for future personalization. + + +## Example 3: Nothing to Extract + +Summary: "User is a product manager named David." +Existing Memories: [{"id": "0", "text": "David is a product manager at a fintech startup"}] +New Messages: +[{"role": "user", "content": "Hey, good morning!"}, + {"role": "assistant", "content": "Good morning, David! How can I help you today?"}] +Observation Date: 2025-08-19 + +Output: {"memory": []} + + +## Example 4: Temporal Precision and Meaning Preservation + +Summary: "" +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "I'm feeling sluggish today — I didn't get to bed until 2 AM last Wednesday, which made Thursday morning a real struggle."}] +Observation Date: 2023-05-25 + +Output: +{"memory": [ + {"id": "0", "text": "User went to bed at 2 AM on Wednesday, May 24, 2023, resulting in a difficult Thursday morning"} +]} + +"Didn't get to bed until 2 AM" means a late BEDTIME of 2 AM. "Last Wednesday" is grounded to May 24 using Observation Date. + +## Example 5: Deduplication — Skip Already Captured + +Recently Extracted: ["Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"] +Existing Memories: [{"id": "0", "text": "Marcus was promoted to Senior Engineer at Shopify around August 12, 2025"}] +New Messages: +[{"role": "user", "content": "Still can't believe I got the senior engineer promotion at Shopify!"}] +Observation Date: 2025-08-19 + +Output: {"memory": []} + + +## Example 6: Extract ALL Dimensions — Don't Miss Secondary Info + +Summary: "User is an aspiring stand-up comedian." +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "As an aspiring stand-up comedian, I'm looking for advice on improving my craft. Can you recommend some stand-up comedy specials on Netflix with strong storytelling like John Mulaney's 'Kid Gorgeous'? I also want to find online resources for comedy writing."}, + {"role": "assistant", "content": "For Netflix specials with great storytelling, check out Hasan Minhaj and Mike Birbiglia. For comedy writing, I'd recommend the Comedy Writing Workbook by Gene Perret and the Kill Tony podcast for studying crowd work."}] +Observation Date: 2023-06-01 + +Output: +{"memory": [ + {"id": "0", "text": "User is an aspiring stand-up comedian seeking to improve their craft through studying comedy specials and writing resources"}, + {"id": "1", "text": "User enjoys watching stand-up comedy specials on Netflix, especially those with strong storytelling like John Mulaney's 'Kid Gorgeous'"}, + {"id": "2", "text": "User was recommended Hasan Minhaj and Mike Birbiglia specials for storytelling, the Comedy Writing Workbook by Gene Perret, and the Kill Tony podcast for studying crowd work"} +]} + +Three dimensions: (1) career aspiration, (2) entertainment viewing preference, (3) specific recommendations. Each extracted separately. + + +## Example 7: Vague Temporal References with Historical Observation Date + +Recently Extracted: ["User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"] +Existing Memories: [{"id": "0", "text": "User started reading 'The Hitchhiker's Guide to the Galaxy' on January 16, 2022"}] +New Messages: +[{"role": "user", "content": "I've actually listened to Ready Player One as an audiobook recently and enjoyed the pop culture references."}] +Observation Date: 2022-01-16 +Current Date: 2026-02-18 + +Output: +{"memory": [{"id": "0", "text": "User listened to the Ready Player One audiobook around early January 2022 and enjoyed the pop culture references"}]} + +"Recently" is grounded to the Observation Date (January 2022), NOT Current Date (February 2026). The Hitchhiker's Guide memory already exists — not re-extracted. + + +## Example 8: Document / Reference Material — Extract Content, Not Actions + +Summary: "" +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "I want you to remember this case. If you understand, just say acknowledged. Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297 — The construction began in 2014, contract signed in 2015 with completion due by October 2015. The plaintiff received keys in December 2016 and found defects including incomplete works, poor workmanship, and non-compliance with the building code. The tribunal found the builder breached contract."}, + {"role": "assistant", "content": "Acknowledged."}] +Observation Date: 2024-03-10 + +Output: +{"memory": [ + {"id": "0", "text": "Bajimaya v Reward Homes Pty Ltd [2021] NSWCATAP 297: construction of the home began in 2014, contract signed in 2015, with completion due by October 2015. Keys were delivered in December 2016."}, + {"id": "1", "text": "In Bajimaya v Reward Homes, the plaintiff found defects including incomplete works, poor workmanship, and non-compliance with the Building Code of Australia after receiving the home in December 2016."}, + {"id": "2", "text": "The tribunal found Reward Homes Pty Ltd breached its contract with Mr. Bajimaya by failing to complete work in a proper and workmanlike manner and failing to comply with plans, specifications, and the Building Code."} +]} + +The user shared reference material to be remembered. Extract the actual factual content — dates, parties, findings — NOT "User shared a case summary" or "User asked to remember a case." + + +## Example 9: Structured Data with Counts and Specifics + +Summary: "" +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "Here are the enemy stat blocks for our D&D campaign: Mummies (4): AC 11, HP 45, Speed 20 ft, with Curse of the Pharaohs (DC 15 Wisdom) and Mummy Rot (DC 15 Constitution). Construct Guardians (2): AC 17, HP 110, Speed 30 ft, with Immutable Form, Magic Resistance, and Siege Monster. Skeletal Warriors (6): AC 12, HP 22, Speed 30 ft, with Undead Fortitude."}, + {"role": "assistant", "content": "Got it! I've noted all the stat blocks. Ready when you want to start the encounter."}] +Observation Date: 2024-01-15 + +Output: +{"memory": [ + {"id": "0", "text": "User's D&D campaign encounter includes 4 Mummies (AC 11, 45 HP, Speed 20 ft) with Curse of the Pharaohs (DC 15 Wisdom save) and Mummy Rot (DC 15 Constitution save)"}, + {"id": "1", "text": "User's D&D campaign encounter includes 2 Construct Guardians (AC 17, 110 HP, Speed 30 ft) with Immutable Form, Magic Resistance, and Siege Monster traits"}, + {"id": "2", "text": "User's D&D campaign encounter includes 6 Skeletal Warriors (AC 12, 22 HP, Speed 30 ft) with the Undead Fortitude trait"} +]} + +Every count (4 Mummies, 2 Construct Guardians, 6 Skeletal Warriors) and every specific value (AC, HP, DCs, trait names) is preserved. Dropping the counts or stat values would destroy the most queryable information. + + +## Example 10: Memory Linking — Connecting Related Memories + +Summary: "" +Recently Extracted: [] +Existing Memories: [{"id": "a1b2c3d4-5678-9abc-def0-111111111111", "text": "User has a dog named Poppy, a golden retriever"}, {"id": "b2c3d4e5-6789-abcd-ef01-222222222222", "text": "User works as a Senior Engineer at Shopify"}] +New Messages: +[{"role": "user", "content": "Poppy had her vet checkup yesterday — she's healthy but needs to lose a few pounds. Also, I'm switching teams at work next month to the payments platform."}] +Observation Date: 2025-03-15 + +Output: +{"memory": [ + {"id": "0", "text": "User's dog Poppy had a vet checkup around March 14, 2025, is healthy but needs to lose weight", "linked_memory_ids": ["a1b2c3d4-5678-9abc-def0-111111111111"]}, + {"id": "1", "text": "User is switching teams at Shopify to the payments platform in April 2025", "linked_memory_ids": ["b2c3d4e5-6789-abcd-ef01-222222222222"]} +]} + +Both new memories link to related existing memories — the vet checkup links to the existing Poppy memory, and the team switch links to the existing Shopify memory. This enables the system to build a graph of related memories. + + +## Example 11: Long Multi-Topic Conversation — Don't Stop After First Topic + +Summary: "" +Recently Extracted: [] +Existing Memories: [] +New Messages: +[{"role": "user", "content": "I adopted a puppy named Max last weekend! He's a beagle mix."}, + {"role": "assistant", "content": "Congratulations! How's he settling in?"}, + {"role": "user", "content": "Great! Oh, and I also started pottery classes on Tuesdays. Made a mug with my daughter's face on it."}, + {"role": "assistant", "content": "Fun! Sounds like a lot going on."}, + {"role": "user", "content": "Yeah — my sister just moved to Portland too. I'm happy but honestly a bit overwhelmed. My boss gave me a promotion to team lead last week as well."}] +Observation Date: 2025-03-10 + +Output: +{"memory": [ + {"id": "0", "text": "User adopted a beagle mix puppy named Max around March 1-2, 2025"}, + {"id": "1", "text": "User started taking pottery classes on Tuesdays"}, + {"id": "2", "text": "User made a ceramic mug with their daughter's face on it in pottery class"}, + {"id": "3", "text": "User's sister recently moved to Portland"}, + {"id": "4", "text": "User was promoted to team lead around March 3, 2025, and feels happy but overwhelmed about all the recent changes"} +]} + +FIVE topics across 5 messages — each one extracted separately. Do not stop after the first topic (the puppy). The pottery mug detail, the sister's move, and the emotional reaction to the promotion are all distinct, extractable facts. + + +## Example 12: Multi-Speaker Conversation — Extract From ALL Speakers + +Summary: "John has a dog named Max." +Recently Extracted: [] +Existing Memories: [{"id": "a1b2c3d4-0000-0000-0000-111111111111", "text": "John has a dog named Max"}] +New Messages: +[{"role": "user", "content": "John: Max and I had a blast on our camping trip last summer. We hiked, swam, and made great memories. It was a really peaceful experience."}, + {"role": "assistant", "content": "Maria: That sounds amazing! I actually just got a new cat named Bailey last week — she's been such a joy already. Camping with pets is so soul-nourishing."}, + {"role": "user", "content": "John: Congrats on Bailey! Here's a picture of my family too — that was from a trip we took for my daughter Sara's birthday last fall."}] +Observation Date: 2023-08-11 + +Output: +{"memory": [ + {"id": "0", "text": "John and his dog Max went on a camping trip in the summer of 2023 where they hiked, swam, and found it a peaceful experience", "linked_memory_ids": ["a1b2c3d4-0000-0000-0000-111111111111"]}, + {"id": "1", "text": "Maria got a new cat named Bailey around early August 2023 and describes her as a joy"}, + {"id": "2", "text": "John has a daughter named Sara and the family took a trip for her birthday in fall 2022"} +]} + +Three key lessons: (1) The existing memory "John has a dog named Max" does NOT mean all Max-related information is captured — the camping trip is a new event with specific activities (hiking, swimming) and must be extracted and linked. (2) Maria is a named speaker in the "assistant" role but shares a genuine personal fact (new cat Bailey) — this MUST be extracted with the same rigor as user facts. Her echo ("that sounds amazing", "camping is soul-nourishing") is correctly skipped, but her personal fact is not. (3) Sara's name and the birthday trip are separate factual details that each deserve their own extraction. + + +# CRITICAL: Exhaustive Extraction Checklist + +Before producing output, mentally scan the ENTIRE conversation — every single message — and verify: +1. Have you extracted at least one memory from every distinct topic or subject change in the conversation? +2. Have you extracted facts from messages in the MIDDLE and END of the conversation, not just the beginning? +3. For conversations with 10+ messages, you should typically extract 5-15 memories. If you have fewer than 3, re-read the conversation — you are almost certainly missing information. +4. Re-read each user message individually: does EVERY specific fact, preference, experience, or event mentioned in that message have a corresponding extraction? If a single message mentions two distinct facts (e.g., an allergy AND a hobby), both must be captured. + +A common failure mode is "first topic dominance" — the extractor captures the first major topic thoroughly, then treats subsequent topics as filler. This is WRONG. Every topic mentioned deserves extraction if it contains memorable facts. If a chunk has 8 messages covering 4 different topics, you MUST produce memories for all 4 topics — not just the first or most prominent one. + + +# OUTPUT FORMAT + +Return ONLY valid JSON parsable by json.loads(). No text, reasoning, explanations, or wrappers. + +## Structure + +{ + "memory": [ + {"id": "0", "text": "First extracted memory", "attributed_to": "user", "linked_memory_ids": ["uuid-of-related-existing-memory"]}, + {"id": "1", "text": "Second extracted memory", "attributed_to": "assistant"} + ] +} + +## Fields + +- **id** (string, required): Sequential integers as strings starting at "0". +- **text** (string, required): A contextually rich, self-contained factual statement (15-80 words). +- **attributed_to** (string, required): Who this memory is about. Use "user" for facts stated by or about the user (preferences, plans, personal facts). Use "assistant" for information provided by the assistant (recommendations, confirmations, plans created, information researched). +- **linked_memory_ids** (array of strings, optional): IDs of Existing Memories that this new memory relates to. Use the exact IDs from the Existing Memories list. Omit or pass [] if no existing memories are related. + +## Rules + +- Extract every piece of memorable information as a separate memory object. +- If nothing is worth extracting, return: {"memory": []} +- No duplicate IDs. Use double quotes. No trailing commas. + +""" + + +AGENT_CONTEXT_SUFFIX = """ + +## Entity Context + +The primary entity is an AI agent. Frame memories from the agent's perspective: +- For user-stated facts, frame as agent knowledge: "Agent was informed that [fact]" or "Agent learned that [fact]" +- For agent actions, use direct statements: "Agent recommended [X]" or "Agent specializes in [domain]" +- For agent configuration or instructions, capture directly: "Agent is configured to [behavior]" + +The attributed_to field should still reflect the original source: "user" for facts the user stated, "assistant" for things the agent said or did. +""" + + +# --------------------------------------------------------------------------- +# V3 Prompt Builder — constructs the user-side prompt for additive extraction +# Ported from platform/backend/shared/core/utils/prompt_builder.py +# --------------------------------------------------------------------------- + +PAST_MESSAGE_TRUNCATION_LIMIT = 300 + + +def _truncate_content(text, limit=PAST_MESSAGE_TRUNCATION_LIMIT): + """Truncate text to limit characters, appending '...' when shortened.""" + if len(text) <= limit: + return text + return text[:limit] + "..." + + +def _format_summary(summary): + """Extract summary text from a string or dict with a 'summary' key.""" + if isinstance(summary, dict): + return summary.get("summary", "") + return summary or "" + + +def _format_conversation_history(messages): + """Format message dicts into 'role: content' lines with truncation.""" + if not messages: + return "" + result = "" + for msg in messages: + role = msg.get("role", "") + content = msg.get("message") or msg.get("content", "") + if role and content: + result += f"{role}: {_truncate_content(content)}\n" + return result + + +def _serialize_memories(memories): + """JSON-serialize a list of memory objects, defaulting to '[]'.""" + return json.dumps(memories or [], ensure_ascii=False) + + +def _format_new_messages(new_messages): + """Pass through if already a string, otherwise JSON-serialize.""" + if isinstance(new_messages, str): + return new_messages + return json.dumps(new_messages or [], ensure_ascii=False) + + +def _resolve_dates(current_date=None, observation_date=None): + """Resolve current and observation dates, defaulting to today.""" + if current_date is None: + current_date = datetime.now(timezone.utc).date().isoformat() + if observation_date is None: + observation_date = current_date + return current_date, observation_date + + +def generate_additive_extraction_prompt( + summary=None, + recently_extracted_memories=None, + existing_memories=None, + new_messages=None, + *, + last_k_messages=None, + current_date=None, + observation_date=None, + custom_instructions=None, + use_input_language=False, +): + """Build the user prompt for additive (ADD-only) extraction with linking. + + Pairs with ADDITIVE_EXTRACTION_PROMPT system prompt. + The LLM will produce only ADD operations, with optional linked_memory_ids. + """ + current_date, observation_date = _resolve_dates(current_date, observation_date) + + sections = [] + sections.append(f"## Summary\n{_format_summary(summary)}") + sections.append(f"## Last k Messages\n{_format_conversation_history(last_k_messages)}") + sections.append(f"## Recently Extracted Memories\n{_serialize_memories(recently_extracted_memories)}") + sections.append(f"## Existing Memories\n{_serialize_memories(existing_memories)}") + sections.append(f"## New Messages\n{_format_new_messages(new_messages)}") + sections.append(f"## Observation Date\n{observation_date}") + sections.append(f"## Current Date\n{current_date}") + + if custom_instructions: + sections.append(f"## Custom Instructions\n{custom_instructions}") + + if use_input_language: + sections.append( + "## Language Requirement\n" + "CRITICAL: Respond in the SAME LANGUAGE and SCRIPT as the input messages.\n" + "1. Match the language of the user's messages exactly — if they write in Korean, extract in Korean; Japanese in Japanese; etc.\n" + "2. Preserve the exact script/alphabet of the input.\n" + "3. Do NOT translate or transliterate into English unless the input is already in English.\n" + "4. Maintain all quality standards (contextual richness, temporal grounding, etc.) regardless of language.\n" + "5. Technical terms, proper nouns, and brand names should be preserved in their original form as used in the input.\n" + "6. If the input mixes languages (e.g., Hinglish), preserve both the mixed language style AND the script.\n" + "7. For Japanese: explicitly resolve omitted subjects using conversation context.\n" + "8. For CJK languages: maintain appropriate formality level from the source text." + ) + + sections.append("# Output:") + return "\n\n".join(sections) diff --git a/mem0/embeddings/azure_openai.py b/mem0/embeddings/azure_openai.py index 547ec0c81..7818de481 100644 --- a/mem0/embeddings/azure_openai.py +++ b/mem0/embeddings/azure_openai.py @@ -53,3 +53,12 @@ class AzureOpenAIEmbedding(EmbeddingBase): """ text = text.replace("\n", " ") 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.""" + 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)] diff --git a/mem0/embeddings/base.py b/mem0/embeddings/base.py index ed328128b..e10959017 100644 --- a/mem0/embeddings/base.py +++ b/mem0/embeddings/base.py @@ -29,3 +29,19 @@ class EmbeddingBase(ABC): list: The embedding vector. """ pass + + def embed_batch(self, texts, memory_action="add"): + """Embed multiple texts. Override in subclasses for native batch support. + + Default implementation calls embed() sequentially for each text. + Subclasses with native batch APIs (e.g., OpenAI) should override + this for better performance. + + Args: + texts: List of text strings to embed. + memory_action: The action context ("add", "search", "update"). + + Returns: + List of embedding vectors (list of floats), one per input text. + """ + return [self.embed(text, memory_action) for text in texts] diff --git a/mem0/embeddings/openai.py b/mem0/embeddings/openai.py index ba5153e6f..956e1003f 100644 --- a/mem0/embeddings/openai.py +++ b/mem0/embeddings/openai.py @@ -47,3 +47,13 @@ class OpenAIEmbedding(EmbeddingBase): .data[0] .embedding ) + + def embed_batch(self, texts, memory_action="add"): + """Embed multiple texts in a single OpenAI API call.""" + texts = [text.replace("\n", " ") for text in texts] + response = self.client.embeddings.create( + input=texts, + model=self.config.model, + dimensions=self.config.embedding_dims, + ) + return [item.embedding for item in sorted(response.data, key=lambda x: x.index)] diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 99bc0b1e5..79ab0503b 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -17,9 +17,14 @@ from pydantic import ValidationError from mem0.configs.base import MemoryConfig, MemoryItem from mem0.configs.enums import MemoryType from mem0.configs.prompts import ( + ADDITIVE_EXTRACTION_PROMPT, + AGENT_CONTEXT_SUFFIX, + generate_additive_extraction_prompt, PROCEDURAL_MEMORY_SYSTEM_PROMPT, get_update_memory_messages, ) +from mem0.utils.lemmatization import lemmatize_for_bm25 +from mem0.utils.entity_extraction import extract_entities_batch from mem0.exceptions import ValidationError as Mem0ValidationError from mem0.memory.base import MemoryBase from mem0.memory.setup import mem0_dir, setup_config @@ -187,15 +192,28 @@ class Memory(MemoryBase): self.db = SQLiteManager(self.config.history_db_path) self.collection_name = self.config.vector_store.config.collection_name self.api_version = self.config.version - + self.custom_instructions = self.config.custom_instructions + # Initialize reranker if configured self.reranker = None if config.reranker: self.reranker = RerankerFactory.create( - config.reranker.provider, + config.reranker.provider, config.reranker.config ) + # 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, + ) + self.enable_graph = False if self.config.graph_store.config: @@ -232,6 +250,75 @@ class Memory(MemoryBase): ) capture_event("mem0.init", self, {"sync_type": "sync"}) + @property + def entity_store(self): + """Lazily initialize entity store on first use.""" + if self._entity_store is None: + entity_config = _safe_deepcopy_config(self.config.vector_store.config) + entity_collection = f"{self.collection_name}_entities" + # Set collection name on the cloned config + if hasattr(entity_config, 'collection_name'): + entity_config.collection_name = entity_collection + elif isinstance(entity_config, dict): + entity_config['collection_name'] = entity_collection + self._entity_store = VectorStoreFactory.create( + self.config.vector_store.provider, entity_config + ) + return self._entity_store + + @staticmethod + def _build_session_scope(filters): + """Build deterministic session scope string from entity IDs.""" + parts = [] + for key in sorted(["user_id", "agent_id", "run_id"]): + val = filters.get(key) + if val: + parts.append(f"{key}={val}") + return "&".join(parts) + + def _upsert_entity(self, entity_text, entity_type, memory_id, filters): + """Upsert an entity into the entity store, linking it to a memory.""" + try: + entity_embedding = self.embedding_model.embed(entity_text, "add") + search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v} + + existing = self.entity_store.search( + query=entity_text, + vectors=entity_embedding, + limit=1, + filters=search_filters, + ) + + if existing and existing[0].score >= 0.95: + # Update existing entity's linked_memory_ids + match = existing[0] + payload = match.payload or {} + linked_ids = payload.get("linked_memory_ids", []) + if memory_id not in linked_ids: + linked_ids.append(memory_id) + payload["linked_memory_ids"] = linked_ids + self.entity_store.update( + vector_id=match.id, + vector=None, + payload=payload, + ) + else: + # Create new entity + entity_id = str(uuid.uuid4()) + entity_payload = { + "data": entity_text, + "entity_type": entity_type, + "linked_memory_ids": [memory_id], + **{k: v for k, v in search_filters.items()}, + } + self.entity_store.insert( + vectors=[entity_embedding], + ids=[entity_id], + payloads=[entity_payload], + ) + except Exception as e: + logger.warning(f"Entity upsert failed for '{entity_text}': {e}") + @classmethod def from_config(cls, config_dict: Dict[str, Any]): try: @@ -289,6 +376,7 @@ class Memory(MemoryBase): infer: bool = True, memory_type: Optional[str] = None, prompt: Optional[str] = None, + observation_date: Optional[str] = None, ): """ Create a new memory. @@ -367,7 +455,7 @@ class Memory(MemoryBase): messages = parse_vision_messages(messages) with concurrent.futures.ThreadPoolExecutor() as executor: - future1 = executor.submit(self._add_to_vector_store, messages, processed_metadata, effective_filters, infer) + future1 = executor.submit(self._add_to_vector_store, messages, processed_metadata, effective_filters, infer, observation_date) future2 = executor.submit(self._add_to_graph, messages, effective_filters) concurrent.futures.wait([future1, future2]) @@ -383,7 +471,7 @@ class Memory(MemoryBase): return {"results": vector_store_result} - def _add_to_vector_store(self, messages, metadata, filters, infer): + def _add_to_vector_store(self, messages, metadata, filters, infer, observation_date=None): if not infer: returned_memories = [] for message_dict in messages: @@ -420,173 +508,191 @@ class Memory(MemoryBase): ) return returned_memories + # === V3 PHASED BATCH PIPELINE === + + # Phase 0: Context gathering + session_scope = self._build_session_scope(filters) + last_messages = self.db.get_last_messages(session_scope, limit=10) parsed_messages = parse_messages(messages) - if self.config.custom_fact_extraction_prompt: - system_prompt = self.config.custom_fact_extraction_prompt - user_prompt = f"Input:\n{parsed_messages}" - else: - # Determine if this should use agent memory extraction based on agent_id presence - # and role types in messages - is_agent_memory = self._should_use_agent_memory_extraction(messages, metadata) - system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages, is_agent_memory) + # 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") + existing_results = self.vector_store.search( + query=parsed_messages, + vectors=query_embedding, + limit=10, + filters=search_filters, + ) - response = self.llm.generate_response( - messages=[ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": user_prompt}, - ], - response_format={"type": "json_object"}, + # Map UUIDs to integers (anti-hallucination) + existing_memories = [] + uuid_mapping = {} + for idx, mem in enumerate(existing_results): + uuid_mapping[str(idx)] = mem.id + existing_memories.append({"id": str(idx), "text": mem.payload.get("data", "")}) + + # Phase 2: LLM extraction (single call) + is_agent_scoped = bool(filters.get("agent_id")) and not filters.get("user_id") + system_prompt = ADDITIVE_EXTRACTION_PROMPT + if is_agent_scoped: + system_prompt += AGENT_CONTEXT_SUFFIX + + custom_instr = self.custom_instructions + if not custom_instr and self.custom_fact_extraction_prompt: + custom_instr = self.custom_fact_extraction_prompt + + user_prompt = generate_additive_extraction_prompt( + existing_memories=existing_memories, + new_messages=parsed_messages, + last_k_messages=last_messages, + custom_instructions=custom_instr, + observation_date=observation_date, ) + try: + response = self.llm.generate_response( + messages=[ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt}, + ], + response_format={"type": "json_object"}, + ) + except Exception as e: + logger.error(f"LLM extraction failed: {e}") + return [] + + # Parse response try: response = remove_code_blocks(response) - if not response.strip(): - new_retrieved_facts = [] + if not response or not response.strip(): + extracted_memories = [] else: try: - # First try direct JSON parsing - new_retrieved_facts = json.loads(response)["facts"] + extracted_memories = json.loads(response).get("memory", []) except json.JSONDecodeError: - # Try extracting JSON from response using built-in function extracted_json = extract_json(response) - new_retrieved_facts = json.loads(extracted_json)["facts"] + extracted_memories = json.loads(extracted_json).get("memory", []) except Exception as e: - logger.error(f"Error in new_retrieved_facts: {e}") - new_retrieved_facts = [] + logger.error(f"Error parsing extraction response: {e}") + extracted_memories = [] - if not new_retrieved_facts: - logger.debug("No new facts retrieved from input. Skipping memory update LLM call.") + if not extracted_memories: + # Save messages even if nothing extracted + self.db.save_messages(messages, session_scope) + return [] - retrieved_old_memory = [] - new_message_embeddings = {} - # Search for existing memories using the provided session identifiers - # Use all available session identifiers for accurate memory retrieval - search_filters = {} - if filters.get("user_id"): - search_filters["user_id"] = filters["user_id"] - if filters.get("agent_id"): - search_filters["agent_id"] = filters["agent_id"] - if filters.get("run_id"): - search_filters["run_id"] = filters["run_id"] - for new_mem in new_retrieved_facts: - messages_embeddings = self.embedding_model.embed(new_mem, "add") - new_message_embeddings[new_mem] = messages_embeddings - existing_memories = self.vector_store.search( - query=new_mem, - vectors=messages_embeddings, - limit=5, - filters=search_filters, - ) - for mem in existing_memories: - retrieved_old_memory.append({"id": mem.id, "text": mem.payload.get("data", "")}) - - unique_data = {} - for item in retrieved_old_memory: - unique_data[item["id"]] = item - retrieved_old_memory = list(unique_data.values()) - logger.info(f"Total existing memories: {len(retrieved_old_memory)}") - - # mapping UUIDs with integers for handling UUID hallucinations - temp_uuid_mapping = {} - for idx, item in enumerate(retrieved_old_memory): - temp_uuid_mapping[str(idx)] = item["id"] - retrieved_old_memory[idx]["id"] = str(idx) - - if new_retrieved_facts: - function_calling_prompt = get_update_memory_messages( - retrieved_old_memory, new_retrieved_facts, self.config.custom_update_memory_prompt - ) - - try: - response: str = self.llm.generate_response( - messages=[{"role": "user", "content": function_calling_prompt}], - response_format={"type": "json_object"}, - ) - except Exception as e: - logger.error(f"Error in new memory actions response: {e}") - response = "" - - try: - if not response or not response.strip(): - logger.warning("Empty response from LLM, no memories to extract") - new_memories_with_actions = {} - else: - response = remove_code_blocks(response) - new_memories_with_actions = json.loads(response) - except Exception as e: - logger.error(f"Invalid JSON response: {e}") - new_memories_with_actions = {} - else: - new_memories_with_actions = {} - - returned_memories = [] + # Phase 3: Batch embed all extracted memory texts + mem_texts = [m.get("text", "") for m in extracted_memories if m.get("text")] try: - for resp in new_memories_with_actions.get("memory", []): - logger.info(resp) + mem_embeddings_list = self.embedding_model.embed_batch(mem_texts, "add") + embed_map = dict(zip(mem_texts, mem_embeddings_list)) + except Exception: + # Fallback: embed individually + embed_map = {} + for text in mem_texts: try: - action_text = resp.get("text") - if not action_text: - logger.info("Skipping memory entry because of empty `text` field.") - continue - - event_type = resp.get("event") - if event_type == "ADD": - memory_id = self._create_memory( - data=action_text, - existing_embeddings=new_message_embeddings, - metadata=deepcopy(metadata), - ) - returned_memories.append({"id": memory_id, "memory": action_text, "event": event_type}) - elif event_type == "UPDATE": - self._update_memory( - memory_id=temp_uuid_mapping[resp.get("id")], - data=action_text, - existing_embeddings=new_message_embeddings, - metadata=deepcopy(metadata), - ) - returned_memories.append( - { - "id": temp_uuid_mapping[resp.get("id")], - "memory": action_text, - "event": event_type, - "previous_memory": resp.get("old_memory"), - } - ) - elif event_type == "DELETE": - self._delete_memory(memory_id=temp_uuid_mapping[resp.get("id")]) - returned_memories.append( - { - "id": temp_uuid_mapping[resp.get("id")], - "memory": action_text, - "event": event_type, - } - ) - elif event_type == "NONE": - # Even if content doesn't need updating, update session IDs if provided - memory_id = temp_uuid_mapping.get(resp.get("id")) - if memory_id and (metadata.get("agent_id") or metadata.get("run_id")): - # Update only the session identifiers, keep content the same - existing_memory = self.vector_store.get(vector_id=memory_id) - updated_metadata = deepcopy(existing_memory.payload) - if metadata.get("agent_id"): - updated_metadata["agent_id"] = metadata["agent_id"] - if metadata.get("run_id"): - updated_metadata["run_id"] = metadata["run_id"] - updated_metadata["updated_at"] = datetime.now(pytz.timezone("US/Pacific")).isoformat() - - self.vector_store.update( - vector_id=memory_id, - vector=None, # Keep same embeddings - payload=updated_metadata, - ) - logger.info(f"Updated session IDs for memory {memory_id}") - else: - logger.info("NOOP for Memory.") + embed_map[text] = self.embedding_model.embed(text, "add") except Exception as e: - logger.error(f"Error processing memory action: {resp}, Error: {e}") + logger.warning(f"Failed to embed memory text: {e}") + + # Phase 4: Per-memory CPU processing + Phase 5: Hash dedup + # Build set of existing hashes for dedup + existing_hashes = set() + for mem in existing_results: + h = mem.payload.get("hash") if hasattr(mem, "payload") and mem.payload else None + if h: + existing_hashes.add(h) + + records = [] # (memory_id, text, embedding, payload) + seen_hashes = set() # dedup within the current batch + for mem in extracted_memories: + text = mem.get("text") + if not text or text not in embed_map: + continue + + mem_hash = hashlib.md5(text.encode()).hexdigest() + if mem_hash in existing_hashes or mem_hash in seen_hashes: + logger.debug(f"Skipping duplicate memory (hash match): {text[:50]}") + continue + seen_hashes.add(mem_hash) + + text_lemmatized = lemmatize_for_bm25(text) + + memory_id = str(uuid.uuid4()) + mem_metadata = deepcopy(metadata) + mem_metadata["data"] = text + mem_metadata["text_lemmatized"] = text_lemmatized + mem_metadata["hash"] = mem_hash + mem_metadata["created_at"] = datetime.now(pytz.timezone("US/Pacific")).isoformat() + if mem.get("attributed_to"): + mem_metadata["attributed_to"] = mem["attributed_to"] + + records.append((memory_id, text, embed_map[text], mem_metadata)) + + if not records: + self.db.save_messages(messages, session_scope) + return [] + + # Phase 6: Batch persist + all_vectors = [r[2] for r in records] + all_ids = [r[0] for r in records] + all_payloads = [r[3] for r in records] + + try: + self.vector_store.insert( + vectors=all_vectors, + ids=all_ids, + payloads=all_payloads, + ) + except Exception: + # Fallback: insert one by one + for mid, vec, pay in zip(all_ids, all_vectors, all_payloads): + try: + self.vector_store.insert(vectors=[vec], ids=[mid], payloads=[pay]) + except Exception as e: + logger.error(f"Failed to insert memory {mid}: {e}") + + # Batch history + history_records = [ + { + "memory_id": r[0], + "old_memory": None, + "new_memory": r[1], + "event": "ADD", + "created_at": r[3].get("created_at"), + "is_deleted": 0, + } + for r in records + ] + try: + self.db.batch_add_history(history_records) + except Exception: + # Fallback: add one by one + for hr in history_records: + try: + self.db.add_history(hr["memory_id"], None, hr["new_memory"], "ADD", created_at=hr.get("created_at")) + except Exception as e: + logger.error(f"Failed to add history for {hr['memory_id']}: {e}") + + # Phase 7: Batch entity linking + try: + all_texts = [r[1] for r in records] + all_entities = extract_entities_batch(all_texts) + for idx, (memory_id, text, embedding, payload) in enumerate(records): + entities = all_entities[idx] if idx < len(all_entities) else [] + for entity_type, entity_text in entities: + self._upsert_entity(entity_text, entity_type, memory_id, search_filters) except Exception as e: - logger.error(f"Error iterating new_memories_with_actions: {e}") + logger.warning(f"Batch entity linking failed: {e}") + + # Phase 8: Save messages + return + self.db.save_messages(messages, session_scope) + + returned_memories = [ + {"id": r[0], "memory": r[1], "event": "ADD"} + for r in records + ] keys, encoded_ids = process_telemetry_filters(filters) capture_event( @@ -630,7 +736,7 @@ class Memory(MemoryBase): "role", ] - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} result_item = MemoryItem( id=memory.id, @@ -731,7 +837,7 @@ class Memory(MemoryBase): "actor_id", "role", ] - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} formatted_memories = [] for mem in actual_memories: @@ -764,7 +870,7 @@ class Memory(MemoryBase): run_id: Optional[str] = None, limit: int = 100, filters: Optional[Dict[str, Any]] = None, - threshold: Optional[float] = None, + threshold: float = 0.1, rerank: bool = True, ): """ @@ -776,7 +882,7 @@ class Memory(MemoryBase): run_id (str, optional): ID of the run to search for. Defaults to None. limit (int, optional): Limit the number of results. Defaults to 100. filters (dict, optional): Legacy filters to apply to the search. Defaults to None. - threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to None. + threshold (float, optional): Minimum score for a memory to be included in the results. Defaults to 0.1. filters (dict, optional): Enhanced metadata filtering with operators: - {"key": "value"} - exact match - {"key": {"eq": "value"}} - equals @@ -951,10 +1057,70 @@ class Memory(MemoryBase): return True return False - def _search_vector_store(self, query, filters, limit, threshold: Optional[float] = None): - embeddings = self.embedding_model.embed(query, "search") - memories = self.vector_store.search(query=query, vectors=embeddings, limit=limit, filters=filters) + 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 + + # Step 1: Preprocess query + query_lemmatized = lemmatize_for_bm25(query) + query_entities = extract_entities(query) + + # Step 2: Embed query + embeddings = self.embedding_model.embed(query, "search") + + # Step 3: Semantic search (over-fetch for scoring pool) + internal_limit = max(limit * 4, 60) + semantic_results = self.vector_store.search( + query=query, vectors=embeddings, limit=internal_limit, filters=filters + ) + + # Step 4: Keyword search (if store supports it) + keyword_results = self.vector_store.keyword_search( + query=query_lemmatized, limit=internal_limit, filters=filters + ) + + # Step 5: Compute BM25 scores from keyword results + bm25_scores = {} + if keyword_results is not None: + midpoint, steepness = get_bm25_params(query, lemmatized=query_lemmatized) + for mem in keyword_results: + mem_id = str(mem.id) if hasattr(mem, 'id') else str(mem.get('id', '')) + raw_score = mem.score if hasattr(mem, 'score') else mem.get('score', 0) + if raw_score and raw_score > 0: + bm25_scores[mem_id] = normalize_bm25(raw_score, midpoint, steepness) + + # Step 6: Compute entity boosts + entity_boosts = {} + if query_entities: + entity_boosts = self._compute_entity_boosts(query_entities, filters) + + # Step 7: Build candidate set from semantic results + # BM25 acts as a boost signal only (not recall-expanding) -- candidates must + # pass the semantic threshold gate, so only semantic results are candidates. + candidates = [] + for mem in semantic_results: + mem_id = str(mem.id) + candidates.append({ + "id": mem_id, + "score": mem.score, + "payload": mem.payload if hasattr(mem, 'payload') else {}, + }) + + # Step 8: Score and rank + scored_results = score_and_rank( + semantic_results=candidates, + bm25_scores=bm25_scores, + entity_boosts=entity_boosts, + threshold=threshold, + top_k=limit, + ) + + # Step 9: Format results promoted_payload_keys = [ "user_id", "agent_id", @@ -962,33 +1128,104 @@ class Memory(MemoryBase): "actor_id", "role", ] - - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} original_memories = [] - for mem in memories: + for scored in scored_results: + payload = scored.get("payload") or {} + + if not payload.get("data"): + continue # Skip candidates with no payload data + memory_item_dict = MemoryItem( - id=mem.id, - memory=mem.payload.get("data", ""), - hash=mem.payload.get("hash"), - created_at=mem.payload.get("created_at"), - updated_at=mem.payload.get("updated_at"), - score=mem.score, + id=scored["id"], + memory=payload.get("data", ""), + hash=payload.get("hash"), + created_at=payload.get("created_at"), + updated_at=payload.get("updated_at"), + score=scored["score"], ).model_dump() + # Add score breakdown to metadata + memory_item_dict["score_breakdown"] = scored.get("score_breakdown", {}) + for key in promoted_payload_keys: - if key in mem.payload: - memory_item_dict[key] = mem.payload[key] + if key in payload: + memory_item_dict[key] = payload[key] - additional_metadata = {k: v for k, v in mem.payload.items() if k not in core_and_promoted_keys} + additional_metadata = {k: v for k, v in payload.items() if k not in core_and_promoted_keys} if additional_metadata: - memory_item_dict["metadata"] = additional_metadata + if "metadata" not in memory_item_dict: + memory_item_dict["metadata"] = {} + memory_item_dict["metadata"].update(additional_metadata) - if threshold is None or mem.score >= threshold: - original_memories.append(memory_item_dict) + original_memories.append(memory_item_dict) return original_memories + def _compute_entity_boosts(self, query_entities, filters): + """Compute per-memory entity boosts from entity store search. + + For each extracted entity from the query: + 1. Embed the entity text + 2. Search the entity store (threshold >= 0.5) + 3. For each matched entity, boost its linked memories + + 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 = [] + for entity_type, entity_text in query_entities[:8]: + key = entity_text.strip().lower() + if key and key not in seen: + seen.add(key) + deduped.append((entity_type, entity_text)) + + if not deduped: + return {} + + search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v} + memory_boosts = {} + + try: + for _, entity_text in deduped: + entity_embedding = self.embedding_model.embed(entity_text, "search") + matches = self.entity_store.search( + query=entity_text, + vectors=entity_embedding, + limit=500, + filters=search_filters, + ) + + for match in matches: + similarity = match.score if hasattr(match, 'score') else 0.0 + if similarity < 0.5: + continue + + payload = match.payload if hasattr(match, 'payload') else {} + linked_memory_ids = payload.get("linked_memory_ids", []) + if not isinstance(linked_memory_ids, list): + continue + + # Spread-attenuated boost: entities linking to many memories get attenuated + num_linked = max(len(linked_memory_ids), 1) + memory_count_weight = 1.0 / (1.0 + 0.001 * ((num_linked - 1) ** 2)) + boost = similarity * ENTITY_BOOST_WEIGHT * memory_count_weight + + for memory_id in linked_memory_ids: + if memory_id: + memory_key = str(memory_id) + memory_boosts[memory_key] = max(memory_boosts.get(memory_key, 0.0), boost) + + except Exception as e: + logger.warning(f"Entity boost computation failed: {e}") + + return memory_boosts + def update(self, memory_id, data): """ Update a memory by ID. @@ -1232,6 +1469,14 @@ class Memory(MemoryBase): self.vector_store = VectorStoreFactory.create( self.config.vector_store.provider, self.config.vector_store.config ) + # Reset entity store if initialized + if self._entity_store is not None: + try: + self._entity_store.reset() + except Exception as e: + logger.warning(f"Failed to reset entity store: {e}") + self._entity_store = None + capture_event("mem0.reset", self, {"sync_type": "sync"}) def chat(self, query): @@ -1254,12 +1499,14 @@ class AsyncMemory(MemoryBase): self.db = SQLiteManager(self.config.history_db_path) self.collection_name = self.config.vector_store.config.collection_name self.api_version = self.config.version - + self.custom_instructions = self.config.custom_instructions + self._entity_store = None + # Initialize reranker if configured self.reranker = None if config.reranker: self.reranker = RerankerFactory.create( - config.reranker.provider, + config.reranker.provider, config.reranker.config ) @@ -1282,6 +1529,21 @@ class AsyncMemory(MemoryBase): capture_event("mem0.init", self, {"sync_type": "async"}) + @property + def entity_store(self): + """Lazily initialize entity store on first use.""" + if self._entity_store is None: + entity_config = _safe_deepcopy_config(self.config.vector_store.config) + entity_collection = f"{self.collection_name}_entities" + if hasattr(entity_config, 'collection_name'): + entity_config.collection_name = entity_collection + elif isinstance(entity_config, dict): + entity_config['collection_name'] = entity_collection + self._entity_store = VectorStoreFactory.create( + self.config.vector_store.provider, entity_config + ) + return self._entity_store + @classmethod async def from_config(cls, config_dict: Dict[str, Any]): try: @@ -1674,7 +1936,7 @@ class AsyncMemory(MemoryBase): "role", ] - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} result_item = MemoryItem( id=memory.id, @@ -1780,7 +2042,7 @@ class AsyncMemory(MemoryBase): "actor_id", "role", ] - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} formatted_memories = [] for mem in actual_memories: @@ -1813,7 +2075,7 @@ class AsyncMemory(MemoryBase): run_id: Optional[str] = None, limit: int = 100, filters: Optional[Dict[str, Any]] = None, - threshold: Optional[float] = None, + threshold: float = 0.1, metadata_filters: Optional[Dict[str, Any]] = None, rerank: bool = True, ): @@ -2007,12 +2269,67 @@ class AsyncMemory(MemoryBase): return True return False - async def _search_vector_store(self, query, filters, limit, threshold: Optional[float] = None): + 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 + + # Step 1: Preprocess query (CPU-bound) + query_lemmatized = await asyncio.to_thread(lemmatize_for_bm25, query) + query_entities = await asyncio.to_thread(extract_entities, query) + + # Step 2: Embed query embeddings = await asyncio.to_thread(self.embedding_model.embed, query, "search") - memories = await asyncio.to_thread( - self.vector_store.search, query=query, vectors=embeddings, limit=limit, filters=filters + + # Step 3: Semantic search (over-fetch) + internal_limit = max(limit * 4, 60) + semantic_results = await asyncio.to_thread( + self.vector_store.search, query=query, vectors=embeddings, limit=internal_limit, filters=filters ) + # Step 4: Keyword search (if store supports it) + keyword_results = await asyncio.to_thread( + self.vector_store.keyword_search, query=query_lemmatized, limit=internal_limit, filters=filters + ) + + # Step 5: Compute BM25 scores + bm25_scores = {} + if keyword_results is not None: + midpoint, steepness = get_bm25_params(query, lemmatized=query_lemmatized) + for mem in keyword_results: + mem_id = str(mem.id) if hasattr(mem, 'id') else str(mem.get('id', '')) + raw_score = mem.score if hasattr(mem, 'score') else mem.get('score', 0) + if raw_score and raw_score > 0: + bm25_scores[mem_id] = normalize_bm25(raw_score, midpoint, steepness) + + # Step 6: Compute entity boosts + entity_boosts = {} + if query_entities: + entity_boosts = await self._compute_entity_boosts_async(query_entities, filters) + + # Step 7: Build candidate set from semantic results + candidates = [] + for mem in semantic_results: + mem_id = str(mem.id) + candidates.append({ + "id": mem_id, + "score": mem.score, + "payload": mem.payload if hasattr(mem, 'payload') else {}, + }) + + # Step 8: Score and rank + scored_results = score_and_rank( + semantic_results=candidates, + bm25_scores=bm25_scores, + entity_boosts=entity_boosts, + threshold=threshold, + top_k=limit, + ) + + # Step 9: Format results promoted_payload_keys = [ "user_id", "agent_id", @@ -2020,33 +2337,92 @@ class AsyncMemory(MemoryBase): "actor_id", "role", ] - - core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", *promoted_payload_keys} + core_and_promoted_keys = {"data", "hash", "created_at", "updated_at", "id", "text_lemmatized", "attributed_to", *promoted_payload_keys} original_memories = [] - for mem in memories: + for scored in scored_results: + payload = scored.get("payload") or {} + if not payload.get("data"): + continue + memory_item_dict = MemoryItem( - id=mem.id, - memory=mem.payload.get("data", ""), - hash=mem.payload.get("hash"), - created_at=mem.payload.get("created_at"), - updated_at=mem.payload.get("updated_at"), - score=mem.score, + id=scored["id"], + memory=payload.get("data", ""), + hash=payload.get("hash"), + created_at=payload.get("created_at"), + updated_at=payload.get("updated_at"), + score=scored["score"], ).model_dump() + memory_item_dict["score_breakdown"] = scored.get("score_breakdown", {}) + for key in promoted_payload_keys: - if key in mem.payload: - memory_item_dict[key] = mem.payload[key] + if key in payload: + memory_item_dict[key] = payload[key] - additional_metadata = {k: v for k, v in mem.payload.items() if k not in core_and_promoted_keys} + additional_metadata = {k: v for k, v in payload.items() if k not in core_and_promoted_keys} if additional_metadata: - memory_item_dict["metadata"] = additional_metadata + if "metadata" not in memory_item_dict: + memory_item_dict["metadata"] = {} + memory_item_dict["metadata"].update(additional_metadata) - if threshold is None or mem.score >= threshold: - original_memories.append(memory_item_dict) + original_memories.append(memory_item_dict) return original_memories + 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]: + key = entity_text.strip().lower() + if key and key not in seen: + seen.add(key) + deduped.append((entity_type, entity_text)) + + if not deduped: + return {} + + search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v} + memory_boosts = {} + + try: + for _, entity_text in deduped: + entity_embedding = await asyncio.to_thread(self.embedding_model.embed, entity_text, "search") + matches = await asyncio.to_thread( + self.entity_store.search, + query=entity_text, + vectors=entity_embedding, + limit=500, + filters=search_filters, + ) + + for match in matches: + similarity = match.score if hasattr(match, 'score') else 0.0 + if similarity < 0.5: + continue + + payload = match.payload if hasattr(match, 'payload') else {} + linked_memory_ids = payload.get("linked_memory_ids", []) + if not isinstance(linked_memory_ids, list): + continue + + num_linked = max(len(linked_memory_ids), 1) + memory_count_weight = 1.0 / (1.0 + 0.001 * ((num_linked - 1) ** 2)) + boost = similarity * ENTITY_BOOST_WEIGHT * memory_count_weight + + for memory_id in linked_memory_ids: + if memory_id: + memory_key = str(memory_id) + memory_boosts[memory_key] = max(memory_boosts.get(memory_key, 0.0), boost) + + except Exception as e: + logger.warning(f"Entity boost computation failed: {e}") + + return memory_boosts + async def update(self, memory_id, data): """ Update a memory by ID asynchronously. diff --git a/mem0/memory/storage.py b/mem0/memory/storage.py index 967dc0c87..6abda7b1c 100644 --- a/mem0/memory/storage.py +++ b/mem0/memory/storage.py @@ -2,6 +2,7 @@ import logging import sqlite3 import threading import uuid +from datetime import datetime, timezone from typing import Any, Dict, List, Optional logger = logging.getLogger(__name__) @@ -14,6 +15,7 @@ class SQLiteManager: self._lock = threading.Lock() self._migrate_history_table() self._create_history_table() + self._create_messages_table() def _migrate_history_table(self) -> None: """ @@ -123,6 +125,28 @@ class SQLiteManager: logger.error(f"Failed to create history table: {e}") raise + def _create_messages_table(self) -> None: + with self._lock: + try: + self.connection.execute("BEGIN") + self.connection.execute( + """ + CREATE TABLE IF NOT EXISTS messages ( + id TEXT PRIMARY KEY, + session_scope TEXT, + role TEXT, + content TEXT, + name TEXT, + created_at DATETIME + ) + """ + ) + self.connection.execute("COMMIT") + except Exception as e: + self.connection.execute("ROLLBACK") + logger.error(f"Failed to create messages table: {e}") + raise + def add_history( self, memory_id: str, @@ -166,6 +190,40 @@ class SQLiteManager: logger.error(f"Failed to add history record: {e}") raise + def batch_add_history(self, records: List[Dict[str, Any]]) -> None: + with self._lock: + try: + self.connection.execute("BEGIN") + self.connection.executemany( + """ + INSERT INTO history ( + id, memory_id, old_memory, new_memory, event, + created_at, updated_at, is_deleted, actor_id, role + ) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + [ + ( + str(uuid.uuid4()), + record.get("memory_id"), + record.get("old_memory"), + record.get("new_memory"), + record.get("event"), + record.get("created_at"), + record.get("updated_at"), + record.get("is_deleted", 0), + record.get("actor_id"), + record.get("role"), + ) + for record in records + ], + ) + self.connection.execute("COMMIT") + except Exception as e: + self.connection.execute("ROLLBACK") + logger.error(f"Failed to batch add history records: {e}") + raise + def get_history(self, memory_id: str) -> List[Dict[str, Any]]: with self._lock: cur = self.connection.execute( @@ -196,18 +254,89 @@ class SQLiteManager: for r in rows ] + def save_messages(self, messages: List[Dict[str, Any]], session_scope: str) -> None: + if not messages: + return + with self._lock: + try: + self.connection.execute("BEGIN") + now = datetime.now(timezone.utc).isoformat() + for message in messages: + self.connection.execute( + """ + INSERT INTO messages (id, session_scope, role, content, name, created_at) + VALUES (?, ?, ?, ?, ?, ?) + """, + ( + str(uuid.uuid4()), + session_scope, + message.get("role"), + message.get("content"), + message.get("name"), + now, + ), + ) + # Evict old messages beyond the most recent 10 for this scope. + # Wrapped in a derived table to force SQLite to materialize the + # ORDER BY before the outer NOT IN evaluates it. + self.connection.execute( + """ + DELETE FROM messages WHERE session_scope = ? AND id NOT IN ( + SELECT id FROM ( + SELECT id FROM messages WHERE session_scope = ? ORDER BY created_at DESC LIMIT 10 + ) + ) + """, + (session_scope, session_scope), + ) + self.connection.execute("COMMIT") + except Exception as e: + self.connection.execute("ROLLBACK") + logger.error(f"Failed to save messages: {e}") + raise + + def get_last_messages(self, session_scope: str, limit: int = 10) -> List[Dict[str, Any]]: + with self._lock: + # Subquery picks the latest N rows (DESC + LIMIT), outer query + # re-sorts them chronologically (ASC) for the caller. + cur = self.connection.execute( + """ + SELECT role, content, name, created_at FROM ( + SELECT role, content, name, created_at + FROM messages + WHERE session_scope = ? + ORDER BY created_at DESC + LIMIT ? + ) ORDER BY created_at ASC + """, + (session_scope, limit), + ) + rows = cur.fetchall() + + return [ + { + "role": r[0], + "content": r[1], + "name": r[2], + "created_at": r[3], + } + for r in rows + ] + def reset(self) -> None: - """Drop and recreate the history table.""" + """Drop and recreate the history and messages tables.""" with self._lock: try: self.connection.execute("BEGIN") self.connection.execute("DROP TABLE IF EXISTS history") + self.connection.execute("DROP TABLE IF EXISTS messages") self.connection.execute("COMMIT") - self._create_history_table() except Exception as e: self.connection.execute("ROLLBACK") - logger.error(f"Failed to reset history table: {e}") + logger.error(f"Failed to reset tables: {e}") raise + self._create_history_table() + self._create_messages_table() def close(self) -> None: if self.connection: diff --git a/mem0/utils/entity_extraction.py b/mem0/utils/entity_extraction.py new file mode 100644 index 000000000..f949c4699 --- /dev/null +++ b/mem0/utils/entity_extraction.py @@ -0,0 +1,357 @@ +""" +Entity extraction from text using spaCy NLP. + +Extracts four types of entities from text: +- **Proper nouns**: Capitalized multi-word sequences (person names, places, brands) +- **Quoted text**: Text in single or double quotes (titles, specific terms) +- **Noun compounds**: Multi-word noun phrases with specific modifiers (e.g., "machine learning") +- **Noun fallback**: Single nouns from circumstantial compound patterns + +Public API: + extract_entities(text: str) -> List[Tuple[str, str]] + +Internal: + _extract_entities_from_doc(doc) -> List[Tuple[str, str]] +""" + +from __future__ import annotations + +import logging +import re +from typing import List, Tuple + +logger = logging.getLogger(__name__) + +# Words that are too generic to be useful as entity heads +_GENERIC_HEADS = { + "thing", "stuff", "way", "time", "experience", "situation", "case", + "fact", "matter", "issue", "idea", "thought", "feeling", "place", + "area", "part", "kind", "type", "sort", "lot", "bit", "day", "year", + "week", "month", "moment", "instance", "example", "technique", + "method", "approach", "process", "step", "tool", "result", "outcome", + "goal", "task", "item", "topic", "scale", "size", "level", "degree", + "amount", "number", "style", "look", "color", "colour", "shape", + "form", "piece", "section", "side", "end", "edge", "surface", "point", +} + +# Modifiers that describe circumstance, not content +_CIRCUMSTANTIAL_MODS = { + "solo", "individual", "team", "group", "joint", "collaborative", + "first", "last", "next", "previous", "final", "initial", "main", "side", +} + +# Adjectives too vague to make a compound entity specific +_NON_SPECIFIC_ADJ = { + "many", "few", "several", "some", "any", "all", "most", "more", + "less", "much", "little", "enough", "various", "numerous", "multiple", + "countless", "great", "good", "bad", "nice", "terrible", "awful", + "awesome", "amazing", "wonderful", "horrible", "excellent", "poor", + "best", "worst", "fine", "okay", "new", "old", "recent", "past", + "future", "current", "previous", "next", "last", "first", "latest", + "early", "late", "former", "modern", "ancient", "big", "small", + "large", "tiny", "huge", "enormous", "long", "short", "tall", "high", + "low", "wide", "narrow", "thick", "thin", "deep", "shallow", + "similar", "different", "same", "other", "another", "such", "certain", + "important", "main", "major", "minor", "key", "primary", "real", + "actual", "true", "whole", "entire", "full", "complete", "total", + "basic", "simple", "interesting", "boring", "exciting", "special", + "particular", "general", "common", "unique", "rare", "typical", + "usual", "normal", "regular", "possible", "likely", "potential", + "available", "necessary", "only", "solo", "individual", "team", + "group", "joint", "collaborative", "final", "initial", "side", +} + +# Generic tail words to strip from compound entities +_GENERIC_ENDINGS = { + "work", "works", "job", "jobs", "task", "tasks", "stuff", "things", + "thing", "info", "information", "details", "data", "content", + "material", "materials", "activities", "activity", "efforts", "effort", + "options", "option", "choices", "choice", "results", "result", + "output", "outputs", "products", "product", "items", "item", +} + +# Capitalized single words that are too generic to be proper nouns +_GENERIC_CAPS = { + "works", "items", "things", "stuff", "resources", "options", "tips", + "ideas", "steps", "ways", "methods", "tools", "features", "benefits", + "examples", "details", "notes", "instructions", "guidelines", + "recommendations", "suggestions", "overview", "summary", "conclusion", + "introduction", "pros", "cons", "advantages", "disadvantages", +} + +# Markdown/formatting markers to skip during extraction +_FORMATTING_MARKERS = {"*", "-", "+", "\u2022", "\u2013", "\u2014", "#", "##", "###", "**", "__"} + + +def _is_sentence_start(tokens: list, idx: int) -> bool: + """Check if a token is at the start of a sentence or after formatting.""" + if idx == 0: + return True + tok = tokens[idx] + if tok.is_sent_start: + return True + prev = tokens[idx - 1].text + return prev in ".!?:" or prev in _FORMATTING_MARKERS or "\n" in prev + + +def _strip_generic_ending(toks: list) -> list: + """Remove generic trailing words from compound token sequences.""" + if len(toks) <= 1: + return toks + last = toks[-1].lemma_.lower() if hasattr(toks[-1], "lemma_") else toks[-1].lower() + return toks[:-1] if last in _GENERIC_ENDINGS and len(toks) > 2 else toks + + +def _lemmatize_compound(toks: list) -> str: + """Join compound tokens, lemmatizing nouns.""" + return " ".join(t.lemma_ if t.pos_ == "NOUN" else t.text for t in toks) + + +def _has_artifacts(txt: str) -> bool: + """Check for formatting artifacts that indicate non-entity text.""" + return any( + [ + "**" in txt or "__" in txt or ":*" in txt, + re.search(r"\s\*\s|\s\*$|^\*\s", txt), + " " in txt or "\n" in txt or "\t" in txt, + len(txt) > 100, + txt.startswith(("\u2022", "-", "+", "\u2013", "\u2014")), + ] + ) + + +def extract_entities(text: str) -> List[Tuple[str, str]]: + """Extract named entities, quoted text, and noun compounds from text. + + This is the public API that accepts a string. It loads the spaCy model + internally and delegates to _extract_entities_from_doc(). + + Args: + text: Input text to extract entities from. + + Returns: + Deduplicated list of (entity_type, entity_text) tuples. + Entity types: PROPER, QUOTED, COMPOUND, NOUN. + Returns empty list if spaCy is unavailable. + """ + from mem0.utils.spacy_models import get_nlp_full + + nlp = get_nlp_full() + if nlp is None: + return [] + + doc = nlp(text) + return _extract_entities_from_doc(doc) + + +def extract_entities_batch(texts: List[str], batch_size: int = 32) -> List[List[Tuple[str, str]]]: + """Extract entities from multiple texts using spaCy's nlp.pipe() for batched NER. + + Uses spaCy's efficient batch processing pipeline instead of calling + nlp() individually per text. Significantly faster for multiple texts. + + Args: + texts: List of input texts to extract entities from. + batch_size: Number of texts to process in each spaCy batch. + + Returns: + List of entity lists, one per input text. Each entity list contains + (entity_type, entity_text) tuples. Returns list of empty lists if + spaCy is unavailable. + """ + if not texts: + return [] + + from mem0.utils.spacy_models import get_nlp_full + + nlp = get_nlp_full() + if nlp is None: + return [[] for _ in texts] + + results = [] + for doc in nlp.pipe(texts, batch_size=batch_size): + results.append(_extract_entities_from_doc(doc)) + return results + + +def _extract_entities_from_doc(doc) -> List[Tuple[str, str]]: + """Extract entities from a spaCy Doc object. + + Ported from platform's shared.core.utils.entity_extraction.extract_entities(). + """ + entities: List[Tuple[str, str]] = [] + text = doc.text + tokens = list(doc) + + # === PROPER NOUN SEQUENCES === + i = 0 + while i < len(tokens): + tok = tokens[i] + if tok.text in _FORMATTING_MARKERS: + i += 1 + continue + is_cap = tok.text and tok.text[0].isupper() + is_label = i + 1 < len(tokens) and tokens[i + 1].text == ":" + + if is_cap and not is_label and tok.pos_ in {"PROPN", "NOUN", "ADJ"}: + seq = [(tok, i)] + j = i + 1 + while j < len(tokens): + t = tokens[j] + if (t.text and t.text[0].isupper()) or t.text.lower() in { + "'s", "of", "the", "in", "and", "for", "at", "is", + }: + seq.append((t, j)) + j += 1 + else: + break + # Strip trailing function words + while seq and seq[-1][0].text.lower() in {"of", "the", "in", "and", "for", "at", "is", "'s"}: + seq.pop() + if seq: + has_mid_cap = any( + not _is_sentence_start(tokens, idx) + for (t, idx) in seq + if t.text[0].isupper() and t.text.lower() not in {"'s", "of", "the", "in", "and", "for", "at", "is"} + ) + if has_mid_cap: + phrase = "".join(t.text_with_ws for (t, idx) in seq).strip() + if len(phrase) > 2: + entities.append(("PROPER", phrase)) + i = j + else: + i += 1 + + # === QUOTED TEXT === + for m in re.finditer(r'"([^"]+)"', text): + if len(m.group(1).strip()) > 2: + entities.append(("QUOTED", m.group(1).strip())) + for m in re.finditer(r"(?:^|[\s\(\[{,;])'([^']+)'(?=[\s\.,;:!?\)\]]|$)", text): + if len(m.group(1).strip()) > 2: + entities.append(("QUOTED", m.group(1).strip())) + + # === NOUN-NOUN COMPOUNDS === + for chunk in doc.noun_chunks: + chunk_tokens = list(chunk) + split_indices: list = [] + poss_splits: list = [] + for idx, tok in enumerate(chunk_tokens): + if tok.dep_ == "case" and tok.text in {"'s", "\u2019s", "'"}: + split_indices.append(idx) + poss_splits.append(idx) + elif tok.pos_ == "PUNCT" and tok.text in {"'", '"', "\u2018", "\u2019", "\u201c", "\u201d"}: + split_indices.append(idx) + + if split_indices: + groups: list = [] + prev = 0 + for split_idx in split_indices: + if split_idx > prev: + groups.append(chunk_tokens[prev:split_idx]) + if split_idx in poss_splits: + next_split = next((s for s in split_indices if s > split_idx), None) + owned = chunk_tokens[split_idx + 1: next_split if next_split else len(chunk_tokens)] + if owned: + first_content = next((t for t in owned if t.pos_ not in {"PUNCT", "PART"}), None) + if not (first_content and first_content.text and first_content.text[0].isupper()): + prev = next_split if next_split else len(chunk_tokens) + continue + prev = split_idx + 1 + if prev < len(chunk_tokens): + groups.append(chunk_tokens[prev:]) + else: + groups = [chunk_tokens] + + for group in groups: + if not group: + continue + head = next((t for t in reversed(group) if t.pos_ in {"NOUN", "PROPN"}), None) + if not head: + continue + head_generic = head.lemma_.lower() in _GENERIC_HEADS + content = [ + t + for t in group + if t.pos_ not in {"DET", "PRON", "PUNCT", "PART", "ADP", "SCONJ", "NUM"} and (t.pos_ == "ADJ" or not t.is_stop) + ] + if not content: + continue + + compound_toks = [t for t in content if t.dep_ == "compound"] + adj_toks = [t for t in content if t.pos_ == "ADJ" or t.dep_ == "amod"] + has_spec_adj = any(t.lemma_.lower() not in _NON_SPECIFIC_ADJ for t in adj_toks) + if head_generic and not has_spec_adj and not compound_toks: + continue + + if compound_toks: + is_circ = any(t.lemma_.lower() in _CIRCUMSTANTIAL_MODS for t in compound_toks) + if is_circ: + val = head.lemma_ if head.pos_ == "NOUN" else head.text + if len(val) > 2: + entities.append(("NOUN", val)) + else: + filtered = _strip_generic_ending( + [t for t in content if not (t.pos_ == "ADJ" and t.lemma_.lower() in _NON_SPECIFIC_ADJ)] + ) + if filtered: + phrase = _lemmatize_compound(filtered) + if len(phrase) > 3 and " " in phrase: + entities.append(("COMPOUND", phrase)) + elif len(content) > 1 and has_spec_adj: + filtered = _strip_generic_ending( + [t for t in content if not ((t.pos_ == "ADJ" or t.dep_ == "amod") and t.lemma_.lower() in _NON_SPECIFIC_ADJ)] + ) + if filtered: + phrase = _lemmatize_compound(filtered) + if len(phrase) > 3 and " " in phrase: + entities.append(("COMPOUND", phrase)) + + # === FALLBACK: Mis-tagged VERB heads === + processed = {e[1].lower() for e in entities if e[0] == "COMPOUND"} + generic_verb_heads = _GENERIC_HEADS | {"find", "buy", "purchase", "sale", "deal", "trip", "visit"} + + def collect_compounds(head): + return [t for t in doc if t.head == head and t.dep_ == "compound"] + + for tok in doc: + if tok.pos_ == "VERB" and tok.dep_ in {"pobj", "dobj", "nsubj"}: + comps = sorted(collect_compounds(tok), key=lambda t: t.i) + if comps: + phrase_toks = comps if tok.lemma_.lower() in generic_verb_heads else comps + [tok] + phrase = " ".join(t.text for t in phrase_toks) + if phrase.lower() not in processed and len(phrase) > 3 and " " in phrase: + entities.append(("COMPOUND", phrase)) + processed.add(phrase.lower()) + + # === DEDUPLICATION & CLEANUP === + seen: set = set() + deduped = [] + for t, e in entities: + k = e.lower().strip() + if k not in seen and len(k) > 2: + seen.add(k) + deduped.append((t, e)) + + cleaned: List[Tuple[str, str]] = [] + for etype, etext in deduped: + txt = re.sub(r"^\*+\s*|\s*\*+$", "", etext.strip()) + txt = re.sub(r"\s*:+$", "", txt) + txt = re.sub(r"^\d+\s*\.\s*", "", txt) + if not txt or len(txt) <= 2 or _has_artifacts(txt): + continue + if etype == "PROPER" and " " not in txt and txt.lower() in _GENERIC_CAPS: + continue + cleaned.append((etype, txt)) + + # Keep best type per entity (PROPER > COMPOUND > QUOTED > NOUN) + type_pri = {"PROPER": 0, "COMPOUND": 1, "QUOTED": 2, "NOUN": 3, "VERB": 4} + best: dict = {} + for t, e in cleaned: + k = e.lower() + if k not in best or type_pri.get(t, 99) < type_pri.get(best[k][0], 99): + best[k] = (t, e) + deduped = list(best.values()) + + # Remove entities that are substrings of longer entities + all_lower = [e[1].lower() for e in deduped] + return [(t, e) for t, e in deduped if not any(e.lower() != o and e.lower() in o for o in all_lower)] diff --git a/mem0/utils/lemmatization.py b/mem0/utils/lemmatization.py new file mode 100644 index 000000000..cab7d87f5 --- /dev/null +++ b/mem0/utils/lemmatization.py @@ -0,0 +1,50 @@ +""" +BM25 lemmatization for consistent keyword matching. + +Uses spaCy's lemmatizer for better handling of: +- Verb forms: attending/attends/attended -> attend +- Comparatives/superlatives: older/oldest -> old +- Plurals: memories -> memory +- Avoids over-stemming: organization != organize + +Also includes original -ing forms alongside lemmas to handle cases +where spaCy's context-dependent lemmatization produces inconsistent +results (e.g., "meeting" as noun vs verb -> different lemmas). +""" + +from __future__ import annotations + +import logging + +logger = logging.getLogger(__name__) + + +def lemmatize_for_bm25(text: str) -> str: + """Lemmatize text for BM25 matching. + + Returns space-joined lemmas for full-text search. Falls back to + the original text if spaCy is unavailable. + """ + from mem0.utils.spacy_models import get_nlp_lemma + + nlp = get_nlp_lemma() + if nlp is None: + return text + + doc = nlp(text.lower()) + tokens = [] + + for token in doc: + if token.is_punct or token.is_stop: + continue + + lemma = token.lemma_ + if lemma.isalnum(): + tokens.append(lemma) + + # Also add original if it ends in -ing and differs from lemma. + # This handles noun/verb ambiguity (meeting/meet, attending/attend). + if token.text.endswith("ing") and token.text != lemma and token.text.isalnum(): + tokens.append(token.text) + + return " ".join(tokens) diff --git a/mem0/utils/scoring.py b/mem0/utils/scoring.py new file mode 100644 index 000000000..1954f4484 --- /dev/null +++ b/mem0/utils/scoring.py @@ -0,0 +1,126 @@ +""" +Scoring utilities for hybrid retrieval. + +Provides: +- **BM25 normalization**: Sigmoid normalization of raw BM25 scores to [0, 1]. +- **BM25 parameter selection**: Query-length-adaptive sigmoid parameters. +- **Additive scoring**: Combined scoring with semantic + BM25 + entity boost. +""" + +from __future__ import annotations + +import math +from typing import Any, Dict, List, Optional + + +def get_bm25_params(query: str, *, lemmatized: Optional[str] = None) -> tuple: + """Get BM25 sigmoid parameters based on query length. + + Longer queries tend to have higher raw BM25 scores, so we adjust + the sigmoid midpoint and steepness accordingly. + + Returns: + (midpoint, steepness) for sigmoid normalization. + """ + if lemmatized is None: + from mem0.utils.lemmatization import lemmatize_for_bm25 + + lemmatized = lemmatize_for_bm25(query) + num_terms = len(lemmatized.split()) if lemmatized else 1 + + if num_terms <= 3: + return 5.0, 0.7 + elif num_terms <= 6: + return 7.0, 0.6 + elif num_terms <= 9: + return 9.0, 0.5 + elif num_terms <= 15: + return 10.0, 0.5 + else: + return 12.0, 0.5 + + +def normalize_bm25(raw_score: float, midpoint: float, steepness: float) -> float: + """Normalize BM25 score to [0, 1] using logistic sigmoid. + + Args: + raw_score: Raw BM25 score (unbounded, typically 0-20+). + midpoint: Score at which sigmoid outputs 0.5. + steepness: Controls how quickly sigmoid transitions. + + Returns: + Normalized score in range [0, 1]. + """ + return 1.0 / (1.0 + math.exp(-steepness * (raw_score - midpoint))) + + +ENTITY_BOOST_WEIGHT = 0.5 + + +def score_and_rank( + semantic_results: List[Dict[str, Any]], + bm25_scores: Dict[str, float], + entity_boosts: Dict[str, float], + threshold: float, + top_k: int, +) -> List[Dict[str, Any]]: + """Score candidates additively and return top-k results. + + For each candidate: + semantic_score is taken from the result's score field. + combined = (semantic + bm25 + entity_boost) / max_possible + + Threshold gates the semantic score BEFORE combining -- candidates + below the threshold are excluded even if BM25/entity would boost them. + + The divisor adapts based on which signals are active: + - Semantic only: max_possible = 1.0 + - Semantic + BM25: max_possible = 2.0 + - Semantic + BM25 + entity: max_possible = 2.5 + - Semantic + entity (no BM25): max_possible = 1.5 + + Returns: + List of scored result dicts sorted by combined score descending. + """ + has_bm25 = bool(bm25_scores) + has_entity = bool(entity_boosts) + + max_possible = 1.0 + if has_bm25: + max_possible += 1.0 + if has_entity: + max_possible += ENTITY_BOOST_WEIGHT + + scored: List[Dict[str, Any]] = [] + + for result in semantic_results: + mem_id = result.get("id") + if mem_id is None: + continue + + semantic_score = result.get("score", 0.0) + if semantic_score < threshold: + continue + + mem_id_str = str(mem_id) + bm25_score = bm25_scores.get(mem_id_str, 0.0) + entity_boost = entity_boosts.get(mem_id_str, 0.0) + + raw_combined = semantic_score + bm25_score + entity_boost + combined = min(raw_combined / max_possible, 1.0) + + scored.append( + { + "id": mem_id_str, + "score": combined, + "score_breakdown": { + "semantic": semantic_score, + "bm25": bm25_score, + "entity_boost": entity_boost, + }, + "payload": result.get("payload"), + } + ) + + scored.sort(key=lambda x: x["score"], reverse=True) + return scored[:top_k] diff --git a/mem0/utils/spacy_models.py b/mem0/utils/spacy_models.py new file mode 100644 index 000000000..f5dfb263b --- /dev/null +++ b/mem0/utils/spacy_models.py @@ -0,0 +1,86 @@ +""" +Shared spaCy model loader. + +Consolidates spaCy model loading into a single module so that +entity_extraction and lemmatization share one instance instead of +each loading their own copy from disk. +""" + +import logging +import threading + +logger = logging.getLogger(__name__) + +_nlp_full = None +_nlp_lemma = None +_load_failed_full = False +_load_failed_lemma = False +_lock = threading.Lock() + + +def _ensure_model_available(): + """Download en_core_web_sm if not already installed. Does not load the model.""" + import spacy + + if not spacy.util.is_package("en_core_web_sm"): + logger.info("Downloading spaCy model en_core_web_sm...") + try: + from spacy.cli import download + + download("en_core_web_sm") + logger.info("spaCy model en_core_web_sm downloaded successfully") + except Exception as e: + raise RuntimeError( + f"Failed to download spaCy model en_core_web_sm: {e}. " + "Please install manually: python -m spacy download en_core_web_sm" + ) from e + + +def get_nlp_full(): + """Return spaCy model with all pipelines (NER, tagger, etc.) for entity extraction.""" + global _nlp_full, _load_failed_full + if _load_failed_full: + return None + if _nlp_full is not None: + return _nlp_full + with _lock: + if _nlp_full is not None: + return _nlp_full + if _load_failed_full: + return None + try: + _ensure_model_available() + import spacy + + _nlp_full = spacy.load("en_core_web_sm") + logger.info("spaCy full model loaded") + except Exception as e: + logger.warning(f"Failed to load spaCy full model: {e}") + _load_failed_full = True + return None + return _nlp_full + + +def get_nlp_lemma(): + """Return spaCy model with only lemmatizer for BM25 text processing.""" + global _nlp_lemma, _load_failed_lemma + if _load_failed_lemma: + return None + if _nlp_lemma is not None: + return _nlp_lemma + with _lock: + if _nlp_lemma is not None: + return _nlp_lemma + if _load_failed_lemma: + return None + try: + _ensure_model_available() + import spacy + + _nlp_lemma = spacy.load("en_core_web_sm", disable=["ner", "parser"]) + logger.info("spaCy lemma model loaded") + except Exception as e: + logger.warning(f"Failed to load spaCy lemma model: {e}") + _load_failed_lemma = True + return None + return _nlp_lemma diff --git a/mem0/vector_stores/azure_ai_search.py b/mem0/vector_stores/azure_ai_search.py index 6165efc6b..edb6fa51e 100644 --- a/mem0/vector_stores/azure_ai_search.py +++ b/mem0/vector_stores/azure_ai_search.py @@ -246,6 +246,34 @@ class AzureAISearch(VectorStoreBase): results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload)) return results + def keyword_search(self, query, limit=5, filters=None): + """Search for memories using keyword/BM25 text matching (no vector queries). + + Args: + query (str): The text query to search for. + limit (int): Maximum number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results with id, score, and payload. + """ + filter_expression = None + if filters: + filter_expression = self._build_filter_expression(filters) + + search_results = self.search_client.search( + search_text=query, + filter=filter_expression, + top=limit, + search_fields=["payload"], + ) + + results = [] + for result in search_results: + payload = json.loads(extract_json(result["payload"])) + results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload)) + return results + def delete(self, vector_id): """ Delete a vector by ID. diff --git a/mem0/vector_stores/azure_mysql.py b/mem0/vector_stores/azure_mysql.py index 2d9ab373b..d95bcd50f 100644 --- a/mem0/vector_stores/azure_mysql.py +++ b/mem0/vector_stores/azure_mysql.py @@ -300,6 +300,61 @@ class AzureMySQL(VectorStoreBase): for r in scored_results ] + def keyword_search(self, query, limit=5, filters=None): + """ + Search for memories using MySQL FULLTEXT search via MATCH() AGAINST(). + + This method attempts to use a FULLTEXT index on the text_lemmatized column. + If the column or index does not exist, it returns None gracefully. + + Args: + query (str): The text query for keyword-based search. + limit (int, optional): Number of results to return. Defaults to 5. + filters (dict, optional): Filters to apply to the search. Defaults to None. + + Returns: + list: Search results in the same format as search(), or None if FULLTEXT + search is not supported on this collection. + """ + try: + filter_conditions = [] + filter_params = [] + + if filters: + for k, v in filters.items(): + filter_conditions.append("JSON_EXTRACT(payload, %s) = %s") + filter_params.extend([f"$.{k}", json.dumps(v)]) + + filter_clause = "" + if filter_conditions: + filter_clause = " AND " + " AND ".join(filter_conditions) + + with self._get_cursor() as cur: + query_sql = f""" + SELECT id, payload, + MATCH(text_lemmatized) AGAINST(%s IN NATURAL LANGUAGE MODE) AS score + FROM `{self.collection_name}` + WHERE MATCH(text_lemmatized) AGAINST(%s IN NATURAL LANGUAGE MODE) + {filter_clause} + ORDER BY score DESC + LIMIT %s + """ + params = [query, query] + filter_params + [limit] + cur.execute(query_sql, params) + results = cur.fetchall() + + return [ + OutputData( + id=r['id'], + score=float(r['score']), + payload=json.loads(r['payload']) if isinstance(r['payload'], str) else r['payload'], + ) + for r in results + ] + except Exception as e: + logger.debug(f"Keyword search not available for collection {self.collection_name}: {e}") + return None + def delete(self, vector_id: str): """ Delete a vector by ID. diff --git a/mem0/vector_stores/baidu.py b/mem0/vector_stores/baidu.py index 2c211abe9..fb08b14e6 100644 --- a/mem0/vector_stores/baidu.py +++ b/mem0/vector_stores/baidu.py @@ -27,6 +27,7 @@ try: VectorIndex, ) from pymochow.model.table import ( + BM25SearchRequest, FloatVector, Partition, Row, @@ -227,6 +228,48 @@ class BaiduDB(VectorStoreBase): return output + def keyword_search(self, query, limit=5, filters=None): + """ + Perform keyword-based search using Baidu Mochow's BM25 search. + + Args: + query (str): The text query to search for. + limit (int, optional): Number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + list: Search results, or None if the table lacks an inverted index. + """ + try: + search_filter = None + if filters: + search_filter = self._create_filter(filters) + + request = BM25SearchRequest( + index_name="data_bm25_idx", + search_text=query, + limit=limit, + filter=search_filter, + ) + + projections = ["id", "metadata"] + res = self._table.bm25_search(request=request, projections=projections) + + output = [] + for row in res.rows: + row_data = row.get("row", {}) + output_data = OutputData( + id=row_data.get("id"), + score=row.get("score", 0.0), + payload=row_data.get("metadata", {}), + ) + output.append(output_data) + + return output + except Exception as e: + logger.error(f"Error during keyword search for query '{query}': {e}") + return None + def delete(self, vector_id): """ Delete a vector by ID. diff --git a/mem0/vector_stores/base.py b/mem0/vector_stores/base.py index 3e22499d7..3c8287547 100644 --- a/mem0/vector_stores/base.py +++ b/mem0/vector_stores/base.py @@ -56,3 +56,20 @@ class VectorStoreBase(ABC): def reset(self): """Reset by delete the collection and recreate it.""" pass + + def keyword_search(self, query: str, limit: int = 5, filters: dict = None): + """Keyword/BM25 full-text search. Returns None if not supported by this store. + + Override in subclasses that support native keyword/BM25 search. + Returns results in the same format as search() -- list of objects with + id, score, and payload attributes. + + Args: + query: The search query text (should be lemmatized for best results). + limit: Maximum number of results to return. + filters: Optional metadata filters (same format as search filters). + + Returns: + List of search results with id, score, payload, or None if not supported. + """ + return None diff --git a/mem0/vector_stores/databricks.py b/mem0/vector_stores/databricks.py index b77058c5f..a96e9cc8b 100644 --- a/mem0/vector_stores/databricks.py +++ b/mem0/vector_stores/databricks.py @@ -484,6 +484,55 @@ class Databricks(VectorStoreBase): logger.error(f"Search failed: {e}") raise + def keyword_search(self, query, limit=5, filters=None): + """ + Search for memories using full-text keyword search. + + Only supported for DELTA_SYNC index type. Returns None for DIRECT_ACCESS indexes. + + Args: + query (str): Search query text. + limit (int): Maximum number of results. Defaults to 5. + filters (dict, optional): Filters to apply. + + Returns: + List[MemoryResult] or None: Search results, or None if index type is DIRECT_ACCESS. + """ + if self.index_type == VectorIndexType.DIRECT_ACCESS: + logger.warning("keyword_search is not supported for DIRECT_ACCESS index type.") + return None + + try: + filters_json = json.dumps(filters) if filters else None + + sdk_results = self.client.vector_search_indexes.query_index( + index_name=self.fully_qualified_index_name, + columns=self.column_names, + query_text=query, + num_results=limit, + query_type="FULL_TEXT", + filters_json=filters_json, + ) + + result_data = sdk_results.result if hasattr(sdk_results, "result") else sdk_results + data_array = result_data.data_array if getattr(result_data, "data_array", None) else [] + + memory_results = [] + for row in data_array: + row_dict = dict(zip(self.column_names, row)) if isinstance(row, (list, tuple)) else row + score = row_dict.get("score") or ( + row[-1] if isinstance(row, (list, tuple)) and len(row) > len(self.column_names) else None + ) + payload = {k: row_dict.get(k) for k in self.column_names} + payload["data"] = payload.get("memory", "") + memory_id = row_dict.get("memory_id") or row_dict.get("id") + memory_results.append(MemoryResult(id=memory_id, score=score, payload=payload)) + return memory_results + + except Exception as e: + logger.error(f"Keyword search failed: {e}") + raise + def delete(self, vector_id): """ Delete a vector by ID from the Delta table. diff --git a/mem0/vector_stores/elasticsearch.py b/mem0/vector_stores/elasticsearch.py index b73eedcdd..93b711e55 100644 --- a/mem0/vector_stores/elasticsearch.py +++ b/mem0/vector_stores/elasticsearch.py @@ -158,6 +158,49 @@ class ElasticsearchDB(VectorStoreBase): return results + def keyword_search(self, query, limit=5, filters=None): + """Search for memories using BM25 keyword matching. + + Args: + query (str): The text query to search for. + limit (int): Maximum number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results with id, score, and payload. + """ + # Build a multi_match query across text fields in metadata + should_clauses = [ + {"match": {"metadata.data": query}}, + {"match": {"metadata.text_lemmatized": query}}, + ] + + bool_query = { + "should": should_clauses, + "minimum_should_match": 1, + } + + if filters: + filter_conditions = [] + for key, value in filters.items(): + filter_conditions.append({"term": {f"metadata.{key}": value}}) + bool_query["filter"] = filter_conditions + + search_query = { + "size": limit, + "query": {"bool": bool_query}, + } + + response = self.client.search(index=self.collection_name, body=search_query) + + results = [] + for hit in response["hits"]["hits"]: + results.append( + OutputData(id=hit["_id"], score=hit["_score"], payload=hit.get("_source", {}).get("metadata", {})) + ) + + return results + def delete(self, vector_id: str) -> None: """Delete a vector by ID.""" self.client.delete(index=self.collection_name, id=vector_id) diff --git a/mem0/vector_stores/milvus.py b/mem0/vector_stores/milvus.py index 09e49a954..ee0539821 100644 --- a/mem0/vector_stores/milvus.py +++ b/mem0/vector_stores/milvus.py @@ -163,6 +163,39 @@ class MilvusDB(VectorStoreBase): result = self._parse_output(data=hits[0]) return result + def keyword_search(self, query, limit=5, filters=None): + """ + Search for memories using BM25-based full-text search via Milvus sparse vector support. + + Milvus 2.5+ supports native BM25 via full-text search with a SPARSE_FLOAT_VECTOR field. + This method attempts to use that capability. If the collection does not have a sparse + field configured, it returns None gracefully. + + Args: + query (str): The text query for keyword-based search. + limit (int, optional): Number of results to return. Defaults to 5. + filters (dict, optional): Filters to apply to the search. Defaults to None. + + Returns: + list: Search results in the same format as search(), or None if sparse search + is not supported on this collection. + """ + try: + query_filter = self._create_filter(filters) if filters else None + hits = self.client.search( + collection_name=self.collection_name, + data=[query], + anns_field="sparse", + limit=limit, + filter=query_filter, + output_fields=["*"], + ) + result = self._parse_output(data=hits[0]) + return result + except Exception as e: + logger.debug(f"Keyword search not available for collection {self.collection_name}: {e}") + return None + def delete(self, vector_id): """ Delete a vector by ID. diff --git a/mem0/vector_stores/mongodb.py b/mem0/vector_stores/mongodb.py index 2bdebf293..e92da83c4 100644 --- a/mem0/vector_stores/mongodb.py +++ b/mem0/vector_stores/mongodb.py @@ -169,6 +169,58 @@ class MongoDB(VectorStoreBase): output = [OutputData(id=str(doc["_id"]), score=doc.get("score"), payload=doc.get("payload")) for doc in results] return output + def keyword_search(self, query, limit=5, filters=None): + """ + Perform keyword-based search using MongoDB Atlas Search. + + Args: + query (str): The text query to search for. + limit (int, optional): Number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results, or None if Atlas Search index is not available. + """ + try: + collection = self.client[self.db_name][self.collection_name] + search_index_name = f"{self.collection_name}_text_search_index" + + pipeline = [ + { + "$search": { + "index": search_index_name, + "text": { + "query": query, + "path": ["payload.data", "payload.text_lemmatized"], + }, + } + }, + {"$set": {"score": {"$meta": "searchScore"}}}, + {"$project": {"embedding": 0}}, + ] + + # Add filter stage if filters are provided + if filters: + filter_conditions = [] + for key, value in filters.items(): + filter_conditions.append({"payload." + key: value}) + if filter_conditions: + pipeline.insert(1, {"$match": {"$and": filter_conditions}}) + + pipeline.append({"$limit": limit}) + + results = list(collection.aggregate(pipeline)) + logger.info(f"Keyword search completed. Found {len(results)} documents.") + + output = [ + OutputData(id=str(doc["_id"]), score=doc.get("score"), payload=doc.get("payload")) + for doc in results + ] + return output + except Exception as e: + logger.error(f"Error during keyword search for query '{query}': {e}") + return None + def delete(self, vector_id: str) -> None: """ Delete a vector by ID. diff --git a/mem0/vector_stores/opensearch.py b/mem0/vector_stores/opensearch.py index deebae91e..0aaaadbf7 100644 --- a/mem0/vector_stores/opensearch.py +++ b/mem0/vector_stores/opensearch.py @@ -182,6 +182,57 @@ class OpenSearchDB(VectorStoreBase): logger.error(f"Error during search: {e}") return [] + def keyword_search(self, query, limit=5, filters=None): + """Search for memories using BM25 keyword matching. + + Args: + query (str): The text query to search for. + limit (int): Maximum number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results with id, score, and payload. + """ + # Build a multi_match query across text fields in payload + should_clauses = [ + {"match": {"payload.data": query}}, + {"match": {"payload.text_lemmatized": query}}, + ] + + bool_query = { + "should": should_clauses, + "minimum_should_match": 1, + } + + # Apply filters consistently with the existing search() method + filter_clauses = [] + if filters: + for key in ["user_id", "run_id", "agent_id"]: + value = filters.get(key) + if value: + filter_clauses.append({"term": {f"payload.{key}.keyword": value}}) + + if filter_clauses: + bool_query["filter"] = filter_clauses + + query_body = { + "size": limit, + "query": {"bool": bool_query}, + } + + try: + response = self.client.search(index=self.collection_name, body=query_body) + + hits = response["hits"]["hits"] + results = [ + OutputData(id=hit["_source"].get("id"), score=hit["_score"], payload=hit["_source"].get("payload", {})) + for hit in hits[:limit] + ] + return results + except Exception as e: + logger.error(f"Error during keyword search: {e}") + return [] + def delete(self, vector_id: str) -> None: """Delete a vector by custom ID.""" # First, find the document by custom ID diff --git a/mem0/vector_stores/pgvector.py b/mem0/vector_stores/pgvector.py index e2d020a66..890ddc0b2 100644 --- a/mem0/vector_stores/pgvector.py +++ b/mem0/vector_stores/pgvector.py @@ -243,6 +243,49 @@ class PGVector(VectorStoreBase): results = cur.fetchall() return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results] + def keyword_search(self, query, limit=5, filters=None): + """ + Search using PostgreSQL full-text search on lemmatized text. + + Args: + query (str): The search query text. + limit (int, optional): Number of results to return. Defaults to 5. + filters (dict, optional): Filters to apply to the search. Defaults to None. + + Returns: + List[OutputData]: Search results ranked by text relevance. + """ + filter_conditions = [] + filter_params = [] + + if filters: + for k, v in filters.items(): + filter_conditions.append("payload->>%s = %s") + filter_params.extend([k, str(v)]) + + filter_clause = "" + if filter_conditions: + filter_clause = "AND " + " AND ".join(filter_conditions) + + try: + with self._get_cursor() as cur: + cur.execute( + f""" + SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', %s)) AS score, payload + FROM {self.collection_name} + WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', %s) + {filter_clause} + ORDER BY score DESC + LIMIT %s + """, + (query, query, *filter_params, limit), + ) + + results = cur.fetchall() + return [OutputData(id=str(r[0]), score=float(r[1]), payload=r[2]) for r in results] + except Exception: + return None + def delete(self, vector_id: str) -> None: """ Delete a vector by ID. diff --git a/mem0/vector_stores/pinecone.py b/mem0/vector_stores/pinecone.py index 08ccf8bc6..e2a7d54d1 100644 --- a/mem0/vector_stores/pinecone.py +++ b/mem0/vector_stores/pinecone.py @@ -241,6 +241,42 @@ class PineconeDB(VectorStoreBase): results = self._parse_output(response.matches) return results + def keyword_search(self, query, limit=5, filters=None): + """ + Search using BM25 sparse vectors for keyword-based retrieval. + + Args: + query (str): The search query text. + limit (int, optional): Number of results to return. Defaults to 5. + filters (dict, optional): Filters to apply to the search. Defaults to None. + + Returns: + List[OutputData]: Search results, or None if hybrid search is not configured. + """ + if not self.hybrid_search or self.sparse_encoder is None: + return None + + try: + filter_dict = self._create_filter(filters) if filters else None + + sparse_vector = self.sparse_encoder.encode_queries(query) + + query_params = { + "sparse_vector": sparse_vector, + "top_k": limit, + "include_metadata": True, + "include_values": False, + } + + if filter_dict: + query_params["filter"] = filter_dict + + response = self.index.query(**query_params, namespace=self.namespace) + + return self._parse_output(response.matches) + except Exception: + return None + def delete(self, vector_id: Union[str, int]): """ Delete a vector by ID. diff --git a/mem0/vector_stores/qdrant.py b/mem0/vector_stores/qdrant.py index 59ee9a92c..173619000 100644 --- a/mem0/vector_stores/qdrant.py +++ b/mem0/vector_stores/qdrant.py @@ -181,6 +181,34 @@ class Qdrant(VectorStoreBase): ) return hits.points + def keyword_search(self, query, limit=5, filters=None): + """ + Search using BM25 sparse vectors for keyword-based retrieval. + + Args: + query (str): The search query text. + limit (int, optional): Number of results to return. Defaults to 5. + filters (dict, optional): Filters to apply to the search. Defaults to None. + + Returns: + list: Search results, or None if the collection does not support sparse vectors. + """ + try: + from qdrant_client import models + + query_filter = self._create_filter(filters) if filters else None + hits = self.client.query_points( + collection_name=self.collection_name, + query=models.Document(text=query, model="Qdrant/bm25"), + using="bm25", + query_filter=query_filter, + limit=limit, + ) + return hits.points + except Exception as e: + logger.debug(f"BM25 keyword search failed (collection may lack sparse vector config): {e}") + return None + def delete(self, vector_id: int): """ Delete a vector by ID. diff --git a/mem0/vector_stores/redis.py b/mem0/vector_stores/redis.py index 7fb1ada9e..770bf6a8e 100644 --- a/mem0/vector_stores/redis.py +++ b/mem0/vector_stores/redis.py @@ -8,7 +8,7 @@ import pytz import redis from redis.commands.search.query import Query from redisvl.index import SearchIndex -from redisvl.query import VectorQuery +from redisvl.query import TextQuery, VectorQuery from redisvl.query.filter import Tag from mem0.memory.utils import extract_json @@ -182,6 +182,60 @@ class RedisDB(VectorStoreBase): for result in results ] + def keyword_search(self, query, limit=5, filters=None): + """ + Search for memories using BM25 keyword search on the memory field. + + Args: + query (str): Search query text. + limit (int): Maximum number of results. Defaults to 5. + filters (dict, optional): Filters to apply (user_id, agent_id, run_id). + + Returns: + List[MemoryResult]: Search results. + """ + filter_expression = None + if filters: + conditions = [Tag(key) == value for key, value in filters.items() if value is not None] + if conditions: + filter_expression = reduce(lambda x, y: x & y, conditions) + + t = TextQuery( + text=query, + text_field_name="memory", + return_fields=["memory_id", "hash", "agent_id", "run_id", "user_id", "memory", "metadata", "created_at"], + filter_expression=filter_expression, + num_results=limit, + ) + + results = self.index.query(t) + + return [ + MemoryResult( + id=result["memory_id"], + score=result.get("text_score", 1.0), + payload={ + "hash": result["hash"], + "data": result["memory"], + "created_at": datetime.fromtimestamp( + int(result["created_at"]), tz=pytz.timezone("US/Pacific") + ).isoformat(timespec="microseconds"), + **( + { + "updated_at": datetime.fromtimestamp( + int(result["updated_at"]), tz=pytz.timezone("US/Pacific") + ).isoformat(timespec="microseconds") + } + if "updated_at" in result + else {} + ), + **{field: result[field] for field in ["agent_id", "run_id", "user_id"] if field in result}, + **{k: v for k, v in json.loads(extract_json(result["metadata"])).items()}, + }, + ) + for result in results + ] + def delete(self, vector_id): self.index.drop_keys(f"{self.schema['index']['prefix']}:{vector_id}") diff --git a/mem0/vector_stores/upstash_vector.py b/mem0/vector_stores/upstash_vector.py index 82dc0f441..4d58a5c84 100644 --- a/mem0/vector_stores/upstash_vector.py +++ b/mem0/vector_stores/upstash_vector.py @@ -149,6 +149,45 @@ class UpstashVector(VectorStoreBase): for res in response ] + def keyword_search(self, query, limit=5, filters=None): + """ + Perform keyword-based search using Upstash's BM25 sparse search. + + Args: + query (str): The text query to search for. + limit (int, optional): Number of results to return. Defaults to 5. + filters (Dict, optional): Filters to apply to the search. + + Returns: + List[OutputData]: Search results, or None if sparse/BM25 search is not supported. + """ + try: + filters_str = ( + " AND ".join([f"{k} = {self._stringify(v)}" for k, v in filters.items()]) + if filters + else None + ) + + response = self.client.query( + data=query, + top_k=limit, + filter=filters_str or "", + include_metadata=True, + namespace=self.collection_name, + ) + + return [ + OutputData( + id=res.id, + score=res.score, + payload=res.metadata, + ) + for res in response + ] + except Exception as e: + logger.error(f"Error during keyword search for query '{query}': {e}") + return None + def delete(self, vector_id: int): """ Delete a vector by ID. diff --git a/mem0/vector_stores/vertex_ai_vector_search.py b/mem0/vector_stores/vertex_ai_vector_search.py index 9e2a9a5c4..5913ccee5 100644 --- a/mem0/vector_stores/vertex_ai_vector_search.py +++ b/mem0/vector_stores/vertex_ai_vector_search.py @@ -274,6 +274,11 @@ class GoogleMatchingEngine(VectorStoreBase): logger.error("Stack trace: %s", traceback.format_exc()) raise + def keyword_search(self, query, limit=5, filters=None): + # Vertex AI hybrid search requires sparse embeddings configuration. + # Not yet supported - requires HybridQuery with sparse encoder setup. + return None + def delete(self, vector_id: Optional[str] = None, ids: Optional[List[str]] = None) -> bool: """ Delete vectors from the Matching Engine index. diff --git a/mem0/vector_stores/weaviate.py b/mem0/vector_stores/weaviate.py index cb1ed6d3a..59ba2af5c 100644 --- a/mem0/vector_stores/weaviate.py +++ b/mem0/vector_stores/weaviate.py @@ -223,6 +223,55 @@ class Weaviate(VectorStoreBase): ) return results + def keyword_search(self, query, limit=5, filters=None): + """ + Search for memories using BM25 keyword search. + + Args: + query (str): Search query text. + limit (int): Maximum number of results. Defaults to 5. + filters (dict, optional): Filters to apply (user_id, agent_id, run_id). + + Returns: + List[OutputData]: Search results. + """ + collection = self.client.collections.get(str(self.collection_name)) + filter_conditions = [] + if filters: + for key, value in filters.items(): + if value and key in ["user_id", "agent_id", "run_id"]: + filter_conditions.append(Filter.by_property(key).equal(value)) + combined_filter = Filter.all_of(filter_conditions) if filter_conditions else None + response = collection.query.bm25( + query=query, + query_properties=["data"], + limit=limit, + filters=combined_filter, + return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"], + return_metadata=MetadataQuery(score=True), + ) + results = [] + for obj in response.objects: + payload = obj.properties.copy() + + for id_field in ["run_id", "agent_id", "user_id"]: + if id_field in payload and payload[id_field] is None: + del payload[id_field] + + payload["id"] = str(obj.uuid).split("'")[0] + if obj.metadata.score is not None: + score = obj.metadata.score + else: + score = 1.0 + results.append( + OutputData( + id=str(obj.uuid), + score=score, + payload=payload, + ) + ) + return results + def delete(self, vector_id): """ Delete a vector by ID. diff --git a/pyproject.toml b/pyproject.toml index 64c5fb40e..25e28c43f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,13 +14,14 @@ license = "Apache-2.0" license-files = ["LICENSE"] requires-python = ">=3.9,<4.0" dependencies = [ - "qdrant-client>=1.9.1", + "qdrant-client>=1.12.0", "pydantic>=2.7.3", "openai>=1.90.0", "posthog>=3.5.0", "pytz>=2024.1", "sqlalchemy>=2.0.31", "protobuf>=5.29.0,<6.0.0", + "spacy>=3.7.0", ] [project.optional-dependencies] @@ -40,7 +41,7 @@ vector_stores = [ "pinecone<=7.3.0", "pinecone-text>=0.10.0", "faiss-cpu>=1.7.4", - "upstash-vector>=0.1.0", + "upstash-vector>=0.6.0", "azure-search-documents>=11.4.0b8", "psycopg>=3.2.8", "psycopg-pool>=3.2.6,<4.0.0", diff --git a/tests/test_main.py b/tests/test_main.py index 2f548e315..b1df33a73 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -86,7 +86,7 @@ def test_add(memory_instance, version, enable_graph): assert result["results"] == [{"memory": "Test memory", "event": "ADD"}] memory_instance._add_to_vector_store.assert_called_once_with( - [{"role": "user", "content": "Test message"}], {"user_id": "test_user"}, {"user_id": "test_user"}, True + [{"role": "user", "content": "Test message"}], {"user_id": "test_user"}, {"user_id": "test_user"}, True, None ) # Remove the conditional assertion for _add_to_graph @@ -129,36 +129,33 @@ def test_search(memory_instance, version, enable_graph): Mock(id="2", payload={"data": "Memory 2", "user_id": "test_user"}, score=0.8), ] memory_instance.vector_store.search = Mock(return_value=mock_memories) + memory_instance.vector_store.keyword_search = Mock(return_value=None) # No BM25 memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3]) memory_instance.graph.search = Mock(return_value=[{"relation": "test_relation"}]) - result = memory_instance.search("test query", user_id="test_user") + with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test query"), \ + patch("mem0.memory.main.extract_entities", return_value=[]): + result = memory_instance.search("test query", user_id="test_user") - if version == "v1.1": - assert "results" in result - assert len(result["results"]) == 2 - assert result["results"][0]["id"] == "1" - assert result["results"][0]["memory"] == "Memory 1" - assert result["results"][0]["user_id"] == "test_user" - assert result["results"][0]["score"] == 0.9 - if enable_graph: - assert "relations" in result - assert result["relations"] == [{"relation": "test_relation"}] - else: - assert "relations" not in result + assert "results" in result + assert len(result["results"]) == 2 + assert result["results"][0]["id"] == "1" + assert result["results"][0]["memory"] == "Memory 1" + assert result["results"][0]["user_id"] == "test_user" + # Score is now combined score (semantic only since no BM25/entity), still 0.9 + assert result["results"][0]["score"] == pytest.approx(0.9) + assert "score_breakdown" in result["results"][0] + + if enable_graph: + assert "relations" in result + assert result["relations"] == [{"relation": "test_relation"}] else: - assert isinstance(result, dict) - assert "results" in result - assert len(result["results"]) == 2 - assert result["results"][0]["id"] == "1" - assert result["results"][0]["memory"] == "Memory 1" - assert result["results"][0]["user_id"] == "test_user" - assert result["results"][0]["score"] == 0.9 + assert "relations" not in result + # Hybrid pipeline over-fetches: max(100*4, 60) = 400 memory_instance.vector_store.search.assert_called_once_with( - query="test query", vectors=[0.1, 0.2, 0.3], limit=100, filters={"user_id": "test_user"} + query="test query", vectors=[0.1, 0.2, 0.3], limit=400, filters={"user_id": "test_user"} ) - memory_instance.embedding_model.embed.assert_called_once_with("test query", "search") if enable_graph: memory_instance.graph.search.assert_called_once_with("test query", {"user_id": "test_user"}, 100) @@ -261,38 +258,26 @@ def test_get_all(memory_instance, version, enable_graph, expected_result): def test_custom_prompts(memory_custom_instance): + """Test that custom_fact_extraction_prompt is passed as custom_instructions in v3 pipeline.""" messages = [{"role": "user", "content": "Test message"}] from mem0.embeddings.mock import MockEmbeddings + # V3 pipeline returns {"memory": [...]} format (single LLM call) memory_custom_instance.llm.generate_response = Mock() - memory_custom_instance.llm.generate_response.return_value = '{"facts": ["fact1", "fact2"]}' + memory_custom_instance.llm.generate_response.return_value = '{"memory": []}' memory_custom_instance.embedding_model = MockEmbeddings() + memory_custom_instance.vector_store.search = Mock(return_value=[]) with patch("mem0.memory.main.parse_messages", return_value="Test message") as mock_parse_messages: - with patch( - "mem0.memory.main.get_update_memory_messages", return_value="custom update memory prompt" - ) as mock_get_update_memory_messages: - memory_custom_instance.add(messages=messages, user_id="test_user") + memory_custom_instance.add(messages=messages, user_id="test_user") - ## custom prompt - ## - mock_parse_messages.assert_called_once_with(messages) + mock_parse_messages.assert_called_once_with(messages) - memory_custom_instance.llm.generate_response.assert_any_call( - messages=[ - {"role": "system", "content": memory_custom_instance.config.custom_fact_extraction_prompt}, - {"role": "user", "content": f"Input:\n{mock_parse_messages.return_value}"}, - ], - response_format={"type": "json_object"}, - ) + # V3 pipeline makes exactly ONE LLM call (not two) + assert memory_custom_instance.llm.generate_response.call_count == 1 - ## custom update memory prompt - ## - mock_get_update_memory_messages.assert_called_once_with( - [], ["fact1", "fact2"], memory_custom_instance.config.custom_update_memory_prompt - ) - - memory_custom_instance.llm.generate_response.assert_any_call( - messages=[{"role": "user", "content": mock_get_update_memory_messages.return_value}], - response_format={"type": "json_object"}, - ) + # The system prompt should be ADDITIVE_EXTRACTION_PROMPT + call_args = memory_custom_instance.llm.generate_response.call_args + llm_messages = call_args[1]["messages"] if "messages" in call_args[1] else call_args[0][0] + assert llm_messages[0]["role"] == "system" + assert "Memory Extractor" in llm_messages[0]["content"] # From ADDITIVE_EXTRACTION_PROMPT diff --git a/tests/utils/__init__.py b/tests/utils/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/utils/test_entity_extraction.py b/tests/utils/test_entity_extraction.py new file mode 100644 index 000000000..267b569a6 --- /dev/null +++ b/tests/utils/test_entity_extraction.py @@ -0,0 +1,102 @@ +import pytest + + +@pytest.fixture(autouse=True) +def _ensure_spacy(): + """Skip tests if spaCy model is not available.""" + try: + import spacy + spacy.load("en_core_web_sm") + except Exception: + pytest.skip("spaCy en_core_web_sm model not available") + + +class TestExtractEntities: + def test_proper_nouns(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("John Smith works at Google on machine learning projects") + entity_texts = [e[1] for e in entities] + entity_types = [e[0] for e in entities] + # Should extract proper nouns + found_proper = any("John" in t or "Google" in t for t in entity_texts) + assert found_proper, f"Expected proper nouns, got {entities}" + + def test_quoted_text(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities('She is reading "The Great Gatsby" this week') + entity_texts = [e[1] for e in entities] + assert any("Great Gatsby" in t for t in entity_texts), f"Expected quoted text, got {entities}" + + def test_compound_nouns(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("The machine learning engineer built a neural network") + entity_texts = [e[1].lower() for e in entities] + has_compound = any("machine" in t and "learning" in t for t in entity_texts) or \ + any("neural" in t and "network" in t for t in entity_texts) + assert has_compound, f"Expected compound nouns, got {entities}" + + def test_empty_string(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("") + assert entities == [] + + def test_no_entities(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("I like things and stuff") + # Generic words should be filtered out + entity_texts = [e[1].lower() for e in entities] + assert "things" not in entity_texts + assert "stuff" not in entity_texts + + def test_deduplication(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("Google is great. I love working at Google.") + google_count = sum(1 for _, t in entities if "Google" in t) + assert google_count <= 1, f"Expected dedup, got {entities}" + + def test_returns_tuples(self): + from mem0.utils.entity_extraction import extract_entities + + entities = extract_entities("John Smith lives in New York City") + for entity in entities: + assert isinstance(entity, tuple) + assert len(entity) == 2 + assert entity[0] in ("PROPER", "QUOTED", "COMPOUND", "NOUN") + assert isinstance(entity[1], str) + + +class TestExtractEntitiesBatch: + def test_batch_processing(self): + from mem0.utils.entity_extraction import extract_entities_batch + + texts = [ + "John works at Google", + "Mary lives in Paris", + "The cat sat on the mat", + ] + results = extract_entities_batch(texts) + assert len(results) == 3 + assert isinstance(results[0], list) + assert isinstance(results[1], list) + assert isinstance(results[2], list) + + def test_empty_input(self): + from mem0.utils.entity_extraction import extract_entities_batch + + assert extract_entities_batch([]) == [] + + def test_consistency_with_single(self): + from mem0.utils.entity_extraction import extract_entities, extract_entities_batch + + text = "John Smith works at Google headquarters" + single = extract_entities(text) + batch = extract_entities_batch([text]) + assert len(batch) == 1 + # Both should extract the same entities + assert set(t for _, t in single) == set(t for _, t in batch[0]) diff --git a/tests/utils/test_lemmatization.py b/tests/utils/test_lemmatization.py new file mode 100644 index 000000000..7373e182a --- /dev/null +++ b/tests/utils/test_lemmatization.py @@ -0,0 +1,67 @@ +import pytest + + +@pytest.fixture(autouse=True) +def _ensure_spacy(): + """Skip tests if spaCy model is not available.""" + try: + import spacy + spacy.load("en_core_web_sm") + except Exception: + pytest.skip("spaCy en_core_web_sm model not available") + + +class TestLemmatizeForBm25: + def test_basic_lemmatization(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("The cats are running quickly") + assert "cat" in result + assert "run" in result or "running" in result + # Stop words and punctuation should be removed + assert "the" not in result.split() + + def test_verb_forms_normalized(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("she attended multiple meetings yesterday") + assert "attend" in result or "attended" in result + assert "meeting" in result # -ing form preserved alongside lemma + # "multiple" is kept (not a spaCy stop word) + + def test_ing_preservation(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("attending the morning meeting") + tokens = result.split() + # Should have both the lemma and the -ing form + assert "attending" in tokens or "attend" in tokens + + def test_empty_string(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("") + assert result == "" + + def test_punctuation_removed(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("Hello, world! How are you?") + assert "," not in result + assert "!" not in result + assert "?" not in result + + def test_lowercased(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("PYTHON Programming LANGUAGE") + for token in result.split(): + assert token == token.lower() + + def test_stop_words_removed(self): + from mem0.utils.lemmatization import lemmatize_for_bm25 + + result = lemmatize_for_bm25("this is a very simple test of the system") + tokens = result.split() + for stop in ["this", "is", "a", "very", "of", "the"]: + assert stop not in tokens diff --git a/tests/utils/test_scoring.py b/tests/utils/test_scoring.py new file mode 100644 index 000000000..edc381100 --- /dev/null +++ b/tests/utils/test_scoring.py @@ -0,0 +1,142 @@ +import pytest + +from mem0.utils.scoring import ( + get_bm25_params, + normalize_bm25, + score_and_rank, + ENTITY_BOOST_WEIGHT, +) + + +class TestGetBm25Params: + def test_short_query(self): + midpoint, steepness = get_bm25_params("hello world", lemmatized="hello world") + assert midpoint == 5.0 + assert steepness == 0.7 + + def test_medium_query(self): + midpoint, steepness = get_bm25_params("x", lemmatized="one two three four five") + assert midpoint == 7.0 + assert steepness == 0.6 + + def test_long_query(self): + words = " ".join(f"word{i}" for i in range(20)) + midpoint, steepness = get_bm25_params("x", lemmatized=words) + assert midpoint == 12.0 + assert steepness == 0.5 + + def test_empty_lemmatized(self): + midpoint, steepness = get_bm25_params("test", lemmatized="") + # Empty string -> 1 term -> short query params + assert midpoint == 5.0 + + +class TestNormalizeBm25: + def test_at_midpoint(self): + score = normalize_bm25(5.0, 5.0, 0.7) + assert abs(score - 0.5) < 0.01 # Should be ~0.5 at midpoint + + def test_high_score(self): + score = normalize_bm25(20.0, 5.0, 0.7) + assert score > 0.99 # Well above midpoint + + def test_low_score(self): + score = normalize_bm25(0.0, 5.0, 0.7) + assert score < 0.05 # Well below midpoint + + def test_range(self): + for raw in [0, 1, 5, 10, 20, 50]: + score = normalize_bm25(float(raw), 5.0, 0.7) + assert 0.0 <= score <= 1.0 + + +class TestScoreAndRank: + def test_semantic_only(self): + results = [ + {"id": "a", "score": 0.9, "payload": {"data": "mem a"}}, + {"id": "b", "score": 0.5, "payload": {"data": "mem b"}}, + ] + scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10) + assert len(scored) == 2 + # With no BM25/entity, max_possible=1.0, so scores stay the same + assert scored[0]["score"] == pytest.approx(0.9) + assert scored[1]["score"] == pytest.approx(0.5) + + def test_semantic_plus_bm25(self): + results = [ + {"id": "a", "score": 0.8, "payload": {"data": "mem a"}}, + {"id": "b", "score": 0.6, "payload": {"data": "mem b"}}, + ] + bm25 = {"a": 0.3, "b": 0.9} + scored = score_and_rank(results, bm25, {}, threshold=0.1, top_k=10) + # max_possible = 2.0 (semantic + bm25) + # a: (0.8 + 0.3) / 2.0 = 0.55 + # b: (0.6 + 0.9) / 2.0 = 0.75 + assert scored[0]["id"] == "b" # b should rank higher due to BM25 + assert scored[0]["score"] == pytest.approx(0.75) + assert scored[1]["id"] == "a" + assert scored[1]["score"] == pytest.approx(0.55) + + def test_all_three_signals(self): + results = [{"id": "a", "score": 0.8, "payload": {"data": "mem a"}}] + bm25 = {"a": 0.6} + entity = {"a": 0.3} + scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10) + # max_possible = 2.5 + expected = (0.8 + 0.6 + 0.3) / 2.5 + assert scored[0]["score"] == pytest.approx(expected) + + def test_threshold_gates_on_semantic(self): + results = [ + {"id": "a", "score": 0.05, "payload": {"data": "mem a"}}, # Below threshold + {"id": "b", "score": 0.5, "payload": {"data": "mem b"}}, + ] + bm25 = {"a": 0.99} # High BM25 shouldn't save it + scored = score_and_rank(results, bm25, {}, threshold=0.1, top_k=10) + assert len(scored) == 1 + assert scored[0]["id"] == "b" + + def test_top_k_limit(self): + results = [{"id": str(i), "score": 0.5, "payload": {}} for i in range(20)] + scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=5) + assert len(scored) == 5 + + def test_score_breakdown_present(self): + results = [{"id": "a", "score": 0.8, "payload": {"data": "x"}}] + bm25 = {"a": 0.4} + entity = {"a": 0.2} + scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10) + breakdown = scored[0]["score_breakdown"] + assert breakdown["semantic"] == 0.8 + assert breakdown["bm25"] == 0.4 + assert breakdown["entity_boost"] == 0.2 + + def test_adaptive_divisor_semantic_only(self): + results = [{"id": "a", "score": 0.8, "payload": {}}] + scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10) + # max_possible = 1.0 (no bm25, no entity) + assert scored[0]["score"] == pytest.approx(0.8) + + def test_adaptive_divisor_semantic_plus_entity(self): + results = [{"id": "a", "score": 0.8, "payload": {}}] + entity = {"a": 0.3} + scored = score_and_rank(results, {}, entity, threshold=0.1, top_k=10) + # max_possible = 1.5 (semantic + entity) + expected = (0.8 + 0.3) / 1.5 + assert scored[0]["score"] == pytest.approx(expected) + + def test_empty_results(self): + scored = score_and_rank([], {}, {}, threshold=0.1, top_k=10) + assert scored == [] + + def test_score_clamped_to_1(self): + results = [{"id": "a", "score": 1.0, "payload": {}}] + bm25 = {"a": 1.0} + entity = {"a": 0.5} + scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10) + assert scored[0]["score"] <= 1.0 + + +class TestEntityBoostWeight: + def test_weight_value(self): + assert ENTITY_BOOST_WEIGHT == 0.5