Compare commits

...

13 Commits

Author SHA1 Message Date
Dev Khant d7a26bd0c3 Add infer param and version bump (#2389)
Co-authored-by: Deshraj Yadav <deshrajdry@gmail.com>
2025-03-18 01:04:58 +05:30
Farzad Sunavala e25dc4b504 bugfix: update Azure AI Search Config (#2380) 2025-03-17 22:19:46 +05:30
Parshva Daftari dab3349990 Neo4j embeddings error (#2377) 2025-03-17 21:57:23 +05:30
Saket Aryan 6db87e8d07 Make DEMO UI Responsive (#2382) 2025-03-14 20:00:28 -07:00
Saket Aryan faf811ee2d Added Custom Categories in Mem0-TS (#2370) 2025-03-14 22:07:46 +05:30
Anusha Yella ee80a43810 Remove tools from LLMs (#2363) 2025-03-14 17:42:48 +05:30
Dev Khant 4be426f762 version bump -> 0.1.68 (#2369) 2025-03-12 21:22:49 +05:30
Farzad Sunavala ba9c61938b feat: enhance Azure AI Search Integration with Binary Quantization, Pre/Post Filter Options, and user agent header (#2354) 2025-03-12 21:20:25 +05:30
Parshva Daftari 65f826e064 Fix langchain neo4j deprecation warning (#2350) 2025-03-12 15:30:45 +05:30
Saket Aryan b43363cdf3 OpenAI Inbuilt Tools (#2362) 2025-03-11 15:33:49 -07:00
Prateek Chhikara 89e786a88e Added agentic tool in docs (#2361) 2025-03-11 13:36:20 -07:00
Saket Aryan 2d5062bd40 Updated Demo (#2360) 2025-03-12 01:49:08 +05:30
Parshva Daftari b89628322d WeaviateDB Integration (#2339) 2025-03-11 00:12:17 +05:30
55 changed files with 2875 additions and 1575 deletions
+1 -1
View File
@@ -12,7 +12,7 @@ install:
install_all:
poetry install
poetry run pip install groq together boto3 litellm ollama chromadb sentence_transformers vertexai \
poetry run pip install groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
google-generativeai elasticsearch opensearch-py vecs
# Format code with ruff
@@ -1,12 +1,14 @@
[Azure AI Search](https://learn.microsoft.com/en-us/azure/search/search-what-is-azure-search/) (formerly known as "Azure Cognitive Search") provides secure information retrieval at scale over user-owned content in traditional and generative AI search applications.
# Azure AI Search
### Usage
[Azure AI Search](https://learn.microsoft.com/azure/search/search-what-is-azure-search/) (formerly known as "Azure Cognitive Search") provides secure information retrieval at scale over user-owned content in traditional and generative AI search applications.
## Usage
```python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "sk-xx" #this key is used for embedding purpose
os.environ["OPENAI_API_KEY"] = "sk-xx" # This key is used for embedding purpose
config = {
"vector_store": {
@@ -15,8 +17,8 @@ config = {
"service_name": "ai-search-test",
"api_key": "*****",
"collection_name": "mem0",
"embedding_model_dims": 1536 ,
"use_compression": False
"embedding_model_dims": 1536,
"compression_type": "none"
}
}
}
@@ -25,20 +27,61 @@ m = Memory.from_config(config)
messages = [
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
]
m.add(messages, user_id="alice", metadata={"category": "movies"})
```
### Config
## Advanced Usage
Let's see the available parameters for the `qdrant` config:
service_name (str): Azure Cognitive Search service name.
| Parameter | Description | Default Value |
| --- | --- | --- |
| `service_name` | Azure AI Search service name | `None` |
| `api_key` | API key of the Azure AI Search service | `None` |
| `collection_name` | The name of the collection/index to store the vectors, it will be created automatically if not exist | `mem0` |
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
| `use_compression` | Use scalar quantization vector compression | False |
```python
# Search with specific filter mode
result = m.search(
"sci-fi movies",
filters={"user_id": "alice"},
limit=5,
vector_filter_mode="preFilter" # Apply filters before vector search
)
# Using binary compression for large vector collections
config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "ai-search-test",
"api_key": "*****",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"compression_type": "binary",
"use_float16": True # Use half precision for storage efficiency
}
}
}
```
## Configuration Parameters
| Parameter | Description | Default Value | Options |
| --- | --- | --- | --- |
| `service_name` | Azure AI Search service name | Required | - |
| `api_key` | API key of the Azure AI Search service | Required | - |
| `collection_name` | The name of the collection/index to store vectors | `mem0` | Any valid index name |
| `embedding_model_dims` | Dimensions of the embedding model | `1536` | Any integer value |
| `compression_type` | Type of vector compression to use | `none` | `none`, `scalar`, `binary` |
| `use_float16` | Store vectors in half precision (Edm.Half) | `False` | `True`, `False` |
## Notes on Configuration Options
- **compression_type**:
- `none`: No compression, uses full vector precision
- `scalar`: Scalar quantization with reasonable balance of speed and accuracy
- `binary`: Binary quantization for maximum compression with some accuracy trade-off
- **vector_filter_mode**:
- `preFilter`: Applies filters before vector search (faster)
- `postFilter`: Applies filters after vector search (may provide better relevance)
- **use_float16**: Using half precision (float16) reduces storage requirements but may slightly impact accuracy. Useful for very large vector collections.
- **Filterable Fields**: The implementation automatically extracts `user_id`, `run_id`, and `agent_id` fields from payloads for filtering.
@@ -0,0 +1,47 @@
[Weaviate](https://weaviate.io/) is an open-source vector search engine. It allows efficient storage and retrieval of high-dimensional vector embeddings, enabling powerful search and retrieval capabilities.
### Installation
```bash
pip install weaviate weaviate-client
```
### Usage
```python Python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "sk-xx"
config = {
"vector_store": {
"provider": "weaviate",
"config": {
"collection_name": "test",
"cluster_url": "http://localhost:8080",
"auth_client_secret": None,
}
}
}
m = Memory.from_config(config)
messages = [
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
{"role": "assistant", "content": "How about a thriller movie? They can be quite engaging."},
{"role": "user", "content": "I’m not a big fan of thriller movies but I love sci-fi movies."},
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
]
m.add(messages, user_id="alice", metadata={"category": "movies"})
```
### Config
Let's see the available parameters for the `weaviate` config:
| Parameter | Description | Default Value |
| --- | --- | --- |
| `collection_name` | The name of the collection to store the vectors | `mem0` |
| `embedding_model_dims` | Dimensions of the embedding model | `1536` |
| `cluster_url` | URL for the Weaviate server | `None` |
| `auth_client_secret` | API key for Weaviate authentication | `None` |
+1
View File
@@ -25,6 +25,7 @@ See the list of supported vector databases below.
<Card title="OpenSearch" href="/components/vectordbs/dbs/opensearch"></Card>
<Card title="Supabase" href="/components/vectordbs/dbs/supabase"></Card>
<Card title="Vertex AI Vector Search" href="/components/vectordbs/dbs/vertex_ai_vector_search"></Card>
<Card title="Weaviate" href="/components/vectordbs/dbs/weaviate"></Card>
</CardGroup>
## Usage
+5 -2
View File
@@ -130,7 +130,8 @@
"components/vectordbs/dbs/elasticsearch",
"components/vectordbs/dbs/opensearch",
"components/vectordbs/dbs/supabase",
"components/vectordbs/dbs/vertex_ai_vector_search"
"components/vectordbs/dbs/vertex_ai_vector_search",
"components/vectordbs/dbs/weaviate"
]
}
]
@@ -185,7 +186,9 @@
"examples/chrome-extension",
"examples/document-writing",
"examples/multimodal-demo",
"examples/personalized-deep-research"
"examples/personalized-deep-research",
"examples/mem0-agentic-tool",
"examples/openai-inbuilt-tools"
]
}
]
+226
View File
@@ -0,0 +1,226 @@
---
title: Mem0 as an Agentic Tool
---
Integrate Mem0's memory capabilities with OpenAI's Agents SDK to create AI agents with persistent memory.
You can create agents that remember past conversations and use that context to provide better responses.
## Installation
First, install the required packages:
```bash
pip install mem0ai pydantic openai-agents
```
You'll also need a custom agents framework for this implementation.
## Setting Up Environment Variables
Store your Mem0 API key as an environment variable:
```bash
export MEM0_API_KEY="your_mem0_api_key"
```
Or in your Python script:
```python
import os
os.environ["MEM0_API_KEY"] = "your_mem0_api_key"
```
## Code Structure
The integration consists of three main components:
1. **Context Manager**: Defines user context for memory operations
2. **Memory Tools**: Functions to add, search, and retrieve memories
3. **Memory Agent**: An agent configured to use these memory tools
## Step-by-Step Implementation
### 1. Import Dependencies
```python
from __future__ import annotations
import os
import asyncio
from pydantic import BaseModel
try:
from mem0 import AsyncMemoryClient
except ImportError:
raise ImportError("mem0 is not installed. Please install it using 'pip install mem0ai'.")
from agents import (
Agent,
ItemHelpers,
MessageOutputItem,
RunContextWrapper,
Runner,
ToolCallItem,
ToolCallOutputItem,
TResponseInputItem,
function_tool,
)
```
### 2. Define Memory Context
```python
class Mem0Context(BaseModel):
user_id: str | None = None
```
### 3. Initialize the Mem0 Client
```python
client = AsyncMemoryClient(api_key=os.getenv("MEM0_API_KEY"))
```
### 4. Create Memory Tools
#### Add to Memory
```python
@function_tool
async def add_to_memory(
context: RunContextWrapper[Mem0Context],
content: str,
) -> str:
"""
Add a message to Mem0
Args:
content: The content to store in memory.
"""
messages = [{"role": "user", "content": content}]
user_id = context.context.user_id or "default_user"
await client.add(messages, user_id=user_id)
return f"Stored message: {content}"
```
#### Search Memory
```python
@function_tool
async def search_memory(
context: RunContextWrapper[Mem0Context],
query: str,
) -> str:
"""
Search for memories in Mem0
Args:
query: The search query.
"""
user_id = context.context.user_id or "default_user"
memories = await client.search(query, user_id=user_id, output_format="v1.1")
results = '\n'.join([result["memory"] for result in memories["results"]])
return str(results)
```
#### Get All Memories
```python
@function_tool
async def get_all_memory(
context: RunContextWrapper[Mem0Context],
) -> str:
"""Retrieve all memories from Mem0"""
user_id = context.context.user_id or "default_user"
memories = await client.get_all(user_id=user_id, output_format="v1.1")
results = '\n'.join([result["memory"] for result in memories["results"]])
return str(results)
```
### 5. Configure the Memory Agent
```python
memory_agent = Agent[Mem0Context](
name="Memory Assistant",
instructions="""You are a helpful assistant with memory capabilities. You can:
1. Store new information using add_to_memory
2. Search existing information using search_memory
3. Retrieve all stored information using get_all_memory
When users ask questions:
- If they want to store information, use add_to_memory
- If they're searching for specific information, use search_memory
- If they want to see everything stored, use get_all_memory""",
tools=[add_to_memory, search_memory, get_all_memory],
)
```
### 6. Implement the Main Runtime Loop
```python
async def main():
current_agent: Agent[Mem0Context] = memory_agent
input_items: list[TResponseInputItem] = []
context = Mem0Context()
while True:
user_input = input("Enter your message (or 'quit' to exit): ")
if user_input.lower() == 'quit':
break
input_items.append({"content": user_input, "role": "user"})
result = await Runner.run(current_agent, input_items, context=context)
for new_item in result.new_items:
agent_name = new_item.agent.name
if isinstance(new_item, MessageOutputItem):
print(f"{agent_name}: {ItemHelpers.text_message_output(new_item)}")
elif isinstance(new_item, ToolCallItem):
print(f"{agent_name}: Calling a tool")
elif isinstance(new_item, ToolCallOutputItem):
print(f"{agent_name}: Tool call output: {new_item.output}")
else:
print(f"{agent_name}: Skipping item: {new_item.__class__.__name__}")
input_items = result.to_input_list()
if __name__ == "__main__":
asyncio.run(main())
```
## Usage Examples
### Storing Information
```
User: Remember that my favorite color is blue
Agent: Calling a tool
Agent: Tool call output: Stored message: my favorite color is blue
Agent: I've stored that your favorite color is blue in my memory. I'll remember that for future conversations.
```
### Searching Memory
```
User: What's my favorite color?
Agent: Calling a tool
Agent: Tool call output: my favorite color is blue
Agent: Your favorite color is blue, based on what you've told me earlier.
```
### Retrieving All Memories
```
User: What do you know about me?
Agent: Calling a tool
Agent: Tool call output: favorite color is blue
my birthday is on March 15
Agent: Based on our previous conversations, I know that:
1. Your favorite color is blue
2. Your birthday is on March 15
```
## Advanced Configuration
### Custom User IDs
You can specify different user IDs to maintain separate memory stores for multiple users:
```python
context = Mem0Context(user_id="user123")
```
## Resources
- [Mem0 Documentation](https://docs.mem0.ai)
- [Mem0 Dashboard](https://app.mem0.ai/dashboard)
- [API Reference](https://docs.mem0.ai/api-reference)
+312
View File
@@ -0,0 +1,312 @@
---
title: OpenAI Inbuilt Tools
---
Integrate Mem0’s memory capabilities with OpenAI’s Inbuilt Tools to create AI agents with persistent memory.
## Getting Started
### Installation
```bash
npm install mem0ai openai zod
```
## Environment Setup
Save your Mem0 and OpenAI API keys in a `.env` file:
```
MEM0_API_KEY=your_mem0_api_key
OPENAI_API_KEY=your_openai_api_key
```
Get your Mem0 API key from the [Mem0 Dashboard](https://app.mem0.ai/dashboard/api-keys).
### Configuration
```javascript
const mem0Config = {
apiKey: process.env.MEM0_API_KEY,
user_id: "sample-user",
};
const openAIClient = new OpenAI();
const mem0Client = new MemoryClient(mem0Config);
```
### Adding Memories
Store user preferences, past interactions, or any relevant information:
<CodeGroup>
```javascript JavaScript
async function addUserPreferences() {
const mem0Client = new MemoryClient(mem0Config);
const userPreferences = "I Love BMW, Audi and Porsche. I Hate Mercedes. I love Red cars and Maroon cars. I have a budget of 120K to 150K USD. I like Audi the most.";
await mem0Client.add([{
role: "user",
content: userPreferences,
}], mem0Config);
}
await addUserPreferences();
```
```json Output (Memories)
[
{
"id": "ff9f3367-9e83-415d-b9c5-dc8befd9a4b4",
"data": { "memory": "Loves BMW, Audi, and Porsche" },
"event": "ADD"
},
{
"id": "04172ce6-3d7b-45a3-b4a1-ee9798593cb4",
"data": { "memory": "Hates Mercedes" },
"event": "ADD"
},
{
"id": "db363a5d-d258-4953-9e4c-777c120de34d",
"data": { "memory": "Loves red cars and maroon cars" },
"event": "ADD"
},
{
"id": "5519aaad-a2ac-4c0d-81d7-0d55c6ecdba8",
"data": { "memory": "Has a budget of 120K to 150K USD" },
"event": "ADD"
},
{
"id": "523b7693-7344-4563-922f-5db08edc8634",
"data": { "memory": "Likes Audi the most" },
"event": "ADD"
}
]
```
</CodeGroup>
### Retrieving Memories
Search for relevant memories based on the current user input:
```javascript
const relevantMemories = await mem0Client.search(userInput, mem0Config);
```
### Structured Responses with Zod
Define structured response schemas to get consistent output formats:
```javascript
// Define the schema for a car recommendation
const CarSchema = z.object({
car_name: z.string(),
car_price: z.string(),
car_url: z.string(),
car_image: z.string(),
car_description: z.string(),
});
// Schema for a list of car recommendations
const Cars = z.object({
cars: z.array(CarSchema),
});
// Create a function tool based on the schema
const carRecommendationTool = zodResponsesFunction({
name: "carRecommendations",
parameters: Cars
});
// Use the tool in your OpenAI request
const response = await openAIClient.responses.create({
model: "gpt-4o",
tools: [{ type: "web_search_preview" }, carRecommendationTool],
input: `${getMemoryString(relevantMemories)}\n${userInput}`,
});
```
### Using Web Search
Combine memory with web search for up-to-date recommendations:
```javascript
const response = await openAIClient.responses.create({
model: "gpt-4o",
tools: [{ type: "web_search_preview" }, carRecommendationTool],
input: `${getMemoryString(relevantMemories)}\n${userInput}`,
});
```
## Examples
### Complete Car Recommendation System
```javascript
import MemoryClient from "mem0ai";
import { OpenAI } from "openai";
import { zodResponsesFunction } from "openai/helpers/zod";
import { z } from "zod";
import dotenv from 'dotenv';
dotenv.config();
const mem0Config = {
apiKey: process.env.MEM0_API_KEY,
user_id: "sample-user",
};
async function run() {
// Responses without memories
console.log("\n\nRESPONSES WITHOUT MEMORIES\n\n");
await main();
// Adding sample memories
await addSampleMemories();
// Responses with memories
console.log("\n\nRESPONSES WITH MEMORIES\n\n");
await main(true);
}
// OpenAI Response Schema
const CarSchema = z.object({
car_name: z.string(),
car_price: z.string(),
car_url: z.string(),
car_image: z.string(),
car_description: z.string(),
});
const Cars = z.object({
cars: z.array(CarSchema),
});
async function main(memory = false) {
const openAIClient = new OpenAI();
const mem0Client = new MemoryClient(mem0Config);
const input = "Suggest me some cars that I can buy today.";
const tool = zodResponsesFunction({ name: "carRecommendations", parameters: Cars });
// Store the user input as a memory
await mem0Client.add([{
role: "user",
content: input,
}], mem0Config);
// Search for relevant memories
let relevantMemories = []
if (memory) {
relevantMemories = await mem0Client.search(input, mem0Config);
}
const response = await openAIClient.responses.create({
model: "gpt-4o",
tools: [{ type: "web_search_preview" }, tool],
input: `${getMemoryString(relevantMemories)}\n${input}`,
});
console.log(response.output);
}
async function addSampleMemories() {
const mem0Client = new MemoryClient(mem0Config);
const myInterests = "I Love BMW, Audi and Porsche. I Hate Mercedes. I love Red cars and Maroon cars. I have a budget of 120K to 150K USD. I like Audi the most.";
await mem0Client.add([{
role: "user",
content: myInterests,
}], mem0Config);
}
const getMemoryString = (memories) => {
const MEMORY_STRING_PREFIX = "These are the memories I have stored. Give more weightage to the question by users and try to answer that first. You have to modify your answer based on the memories I have provided. If the memories are irrelevant you can ignore them. Also don't reply to this section of the prompt, or the memories, they are only for your reference. The MEMORIES of the USER are: \n\n";
const memoryString = memories.map((mem) => `${mem.memory}`).join("\n") ?? "";
return memoryString.length > 0 ? `${MEMORY_STRING_PREFIX}${memoryString}` : "";
};
run().catch(console.error);
```
### Responses
<CodeGroup>
```json Without Memories
{
"cars": [
{
"car_name": "Toyota Camry",
"car_price": "$25,000",
"car_url": "https://www.toyota.com/camry/",
"car_image": "https://link-to-toyota-camry-image.com",
"car_description": "Reliable mid-size sedan with great fuel efficiency."
},
{
"car_name": "Honda Accord",
"car_price": "$26,000",
"car_url": "https://www.honda.com/accord/",
"car_image": "https://link-to-honda-accord-image.com",
"car_description": "Comfortable and spacious with advanced safety features."
},
{
"car_name": "Ford Mustang",
"car_price": "$28,000",
"car_url": "https://www.ford.com/mustang/",
"car_image": "https://link-to-ford-mustang-image.com",
"car_description": "Iconic sports car with powerful engine options."
},
{
"car_name": "Tesla Model 3",
"car_price": "$38,000",
"car_url": "https://www.tesla.com/model3",
"car_image": "https://link-to-tesla-model3-image.com",
"car_description": "Electric vehicle with advanced technology and long range."
},
{
"car_name": "Chevrolet Equinox",
"car_price": "$24,000",
"car_url": "https://www.chevrolet.com/equinox/",
"car_image": "https://link-to-chevron-equinox-image.com",
"car_description": "Compact SUV with a spacious interior and user-friendly technology."
}
]
}
```
```json With Memories
{
"cars": [
{
"car_name": "Audi RS7",
"car_price": "$118,500",
"car_url": "https://www.audiusa.com/us/web/en/models/rs7/2023/overview.html",
"car_image": "https://www.audiusa.com/content/dam/nemo/us/models/rs7/my23/gallery/1920x1080_AOZ_A717_191004.jpg",
"car_description": "The Audi RS7 is a high-performance hatchback with a sleek design, powerful 591-hp twin-turbo V8, and luxurious interior. It's available in various colors including red."
},
{
"car_name": "Porsche Panamera GTS",
"car_price": "$129,300",
"car_url": "https://www.porsche.com/usa/models/panamera/panamera-models/panamera-gts/",
"car_image": "https://files.porsche.com/filestore/image/multimedia/noneporsche-panamera-gts-sample-m02-high/normal/8a6327c3-6c7f-4c6f-a9a8-fb9f58b21795;sP;twebp/porsche-normal.webp",
"car_description": "The Porsche Panamera GTS is a luxury sports sedan with a 473-hp V8 engine, exquisite handling, and available in stunning red. Balances sportiness and comfort."
},
{
"car_name": "BMW M5",
"car_price": "$105,500",
"car_url": "https://www.bmwusa.com/vehicles/m-models/m5/sedan/overview.html",
"car_image": "https://www.bmwusa.com/content/dam/bmwusa/M/m5/2023/bmw-my23-m5-sapphire-black-twilight-purple-exterior-02.jpg",
"car_description": "The BMW M5 is a powerhouse sedan with a 600-hp V8 engine, known for its great handling and luxury. It comes in several distinctive colors including maroon."
}
]
}
```
</CodeGroup>
## Resources
- [Mem0 Documentation](https://docs.mem0.ai)
- [Mem0 Dashboard](https://app.mem0.ai/dashboard)
- [API Reference](https://docs.mem0.ai/api-reference)
- [OpenAI Documentation](https://platform.openai.com/docs)
+8
View File
@@ -58,4 +58,12 @@ Explore how **Mem0** can power real-world applications and bring personalized, i
<Card title="Personalized Research Agent" icon="robot" href="/examples/personalized-deep-research">
Build a **Deep Research AI** that remembers your research goals and compiles insights from vast information sources.
</Card>
<Card title="Mem0 as an Agentic Tool" icon="robot" href="/examples/mem0-agentic-tool">
Integrate Mem0's memory capabilities with OpenAI's Agents SDK to create AI agents with persistent memory.
</Card>
<Card title="OpenAI Inbuilt Tools" icon="robot" href="/examples/openai-inbuilt-tools">
Use Mem0's memory capabilities with OpenAI's Inbuilt Tools to create AI agents with persistent memory.
</Card>
</CardGroup>
+5 -1
View File
@@ -85,14 +85,18 @@ m = Memory.from_config(config_dict=config)
<CodeGroup>
```python Code
const messages = [
messages = [
{"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
{"role": "assistant", "content": "How about a thriller movies? They can be quite engaging."},
{"role": "user", "content": "I'm not a big fan of thriller movies but I love sci-fi movies."},
{"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."}
]
# Store inferred memories (default behavior)
result = m.add(messages, user_id="alice", metadata={"category": "movie_recommendations"})
# Store raw messages without inference
# result = m.add(messages, user_id="alice", metadata={"category": "movie_recommendations"}, infer=False)
```
```json Output
+15 -10
View File
@@ -6,7 +6,7 @@ import { Thread } from "@/components/assistant-ui/thread";
import { ThreadList } from "@/components/assistant-ui/thread-list";
import { useEffect, useState } from "react";
import { v4 as uuidv4 } from "uuid";
import { Sun, Moon, MessageSquare } from "lucide-react";
import { Sun, Moon, AlignJustify } from "lucide-react";
import { Button } from "@/components/ui/button";
import ThemeAwareLogo from "@/components/mem0/theme-aware-logo";
import Link from "next/link";
@@ -63,7 +63,7 @@ export const Assistant = () => {
return (
<AssistantRuntimeProvider runtime={runtime}>
<div className={`h-dvh bg-[#f8fafc] dark:bg-zinc-900 text-[#1e293b] ${isDarkMode ? "dark" : ""}`}>
<div className={`bg-[#f8fafc] dark:bg-zinc-900 text-[#1e293b] ${isDarkMode ? "dark" : ""}`}>
<header className="h-16 border-b border-[#e2e8f0] flex items-center justify-between px-4 sm:px-6 bg-white dark:bg-zinc-900 dark:border-zinc-800 dark:text-white">
<div className="flex items-center">
<Link href="/" className="flex items-center">
@@ -71,29 +71,34 @@ export const Assistant = () => {
</Link>
</div>
<div className="flex items-center">
<Button
<Button
variant="ghost"
size="sm"
onClick={() => setSidebarOpen(true)}
className="text-[#475569] dark:text-zinc-300 md:hidden"
>
<MessageSquare className="w-10 h-10" />
</Button>
<AlignJustify size={24} className="md:hidden" />
</Button>
<div className="md:flex items-center hidden">
<button
className="p-2 rounded-full hover:bg-[#eef2ff] dark:hover:bg-zinc-800 text-[#475569] dark:text-zinc-300"
onClick={toggleDarkMode}
aria-label="Toggle theme"
>
{isDarkMode ? <Sun className="w-5 h-5" /> : <Moon className="w-5 h-5" />}
{isDarkMode ? <Sun className="w-6 h-6" /> : <Moon className="w-6 h-6" />}
</button>
<GithubButton url="https://github.com/mem0ai/mem0/tree/main/examples" />
<Link href={"https://app.mem0.ai/"} target="_blank" className="py-2 ml-2 px-4 font-semibold dark:bg-zinc-100 dark:hover:bg-zinc-200 bg-zinc-800 text-white rounded-full hover:bg-zinc-900 dark:text-[#475569]">
Save Memories
</Link>
</div>
</header>
<div className="grid grid-cols-1 md:grid-cols-[260px_1fr] gap-x-0 h-[calc(100vh-8rem)] md:h-[calc(100vh-4rem)]">
<div className="grid grid-cols-1 md:grid-cols-[260px_1fr] gap-x-0 h-[calc(100dvh-4rem)]">
<ThreadList onResetUserId={resetUserId} isDarkMode={isDarkMode} />
<Thread sidebarOpen={sidebarOpen} setSidebarOpen={setSidebarOpen} onResetUserId={resetUserId} isDarkMode={isDarkMode} />
<Thread sidebarOpen={sidebarOpen} setSidebarOpen={setSidebarOpen} onResetUserId={resetUserId} isDarkMode={isDarkMode} toggleDarkMode={toggleDarkMode} />
</div>
</div>
</AssistantRuntimeProvider>
@@ -19,14 +19,14 @@ import {
AlertDialogTitle,
AlertDialogTrigger,
} from "@/components/ui/alert-dialog";
import ThemeAwareLogo from "@/components/assistant-ui/theme-aware-logo";
import Link from "next/link";
// import ThemeAwareLogo from "@/components/assistant-ui/theme-aware-logo";
// import Link from "next/link";
interface ThreadListProps {
onResetUserId?: () => void;
isDarkMode: boolean;
}
export const ThreadList: FC<ThreadListProps> = ({ onResetUserId, isDarkMode }) => {
export const ThreadList: FC<ThreadListProps> = ({ onResetUserId }) => {
const [open, setOpen] = useState(false);
return (
@@ -79,13 +79,7 @@ export const ThreadList: FC<ThreadListProps> = ({ onResetUserId, isDarkMode }) =
</div>
<ThreadListItems />
</div>
<div>
<Link href="https://www.assistant-ui.com/" target="_blank" className="flex justify-center items-center gap-2">
<h1 className="text-sm text-[#475569] dark:text-zinc-300 text-center">built using</h1>
<ThemeAwareLogo width={24} height={24} isDarkMode={isDarkMode} />
<p className="text-md font-bold dark:text-zinc-300">assistant-ui</p>
</Link>
</div>
</ThreadListPrimitive.Root>
</div>
);
@@ -22,9 +22,12 @@ import {
SendHorizontalIcon,
ArchiveIcon,
PlusIcon,
Sun,
Moon,
SaveIcon,
} from "lucide-react";
import { cn } from "@/lib/utils";
import { Dispatch, SetStateAction, useState } from "react";
import { Dispatch, SetStateAction, useState, useRef } from "react";
import { Button } from "@/components/ui/button";
import { ScrollArea } from "../ui/scroll-area";
import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button";
@@ -42,14 +45,14 @@ import {
AlertDialogTitle,
AlertDialogTrigger,
} from "@/components/ui/alert-dialog";
import GithubButton from "../mem0/github-button";
import Link from "next/link";
import ThemeAwareLogo from "./theme-aware-logo";
interface ThreadProps {
sidebarOpen: boolean;
setSidebarOpen: Dispatch<SetStateAction<boolean>>;
onResetUserId?: () => void;
isDarkMode: boolean;
toggleDarkMode: () => void;
}
export const Thread: FC<ThreadProps> = ({
@@ -57,12 +60,14 @@ export const Thread: FC<ThreadProps> = ({
setSidebarOpen,
onResetUserId,
isDarkMode,
toggleDarkMode
}) => {
const [resetDialogOpen, setResetDialogOpen] = useState(false);
const composerInputRef = useRef<HTMLTextAreaElement>(null);
return (
<ThreadPrimitive.Root
className="bg-[#f8fafc] dark:bg-zinc-900 box-border h-full flex flex-col overflow-hidden relative"
className="bg-[#f8fafc] dark:bg-zinc-900 box-border flex flex-col overflow-hidden relative h-[calc(100dvh-4rem)] pb-4 md:h-full"
style={{
["--thread-max-width" as string]: "42rem",
}}
@@ -78,13 +83,13 @@ export const Thread: FC<ThreadProps> = ({
{/* Mobile sidebar drawer */}
<div
className={cn(
"fixed inset-y-0 left-0 z-40 w-[85%] bg-white dark:bg-zinc-900 transform transition-transform duration-300 ease-in-out md:hidden",
"fixed inset-y-0 left-0 z-40 w-[75%] bg-white shadow-lg rounded-r-lg dark:bg-zinc-900 transform transition-transform duration-300 ease-in-out md:hidden",
sidebarOpen ? "translate-x-0" : "-translate-x-full"
)}
>
<div className="h-full flex flex-col">
<div className="flex items-center justify-between border-b dark:text-white border-[#e2e8f0] dark:border-zinc-800 p-4">
<h2 className="font-medium">Recent Chats</h2>
<h2 className="font-medium">Settings</h2>
<div className="flex items-center gap-2">
{onResetUserId && (
<AlertDialog
@@ -141,13 +146,44 @@ export const Thread: FC<ThreadProps> = ({
<div className="flex flex-col justify-between items-stretch gap-1.5 h-full dark:text-white">
<ThreadListPrimitive.Root className="flex flex-col items-stretch gap-1.5 h-full dark:text-white">
<ThreadListPrimitive.New asChild>
<div className="flex items-center flex-col gap-2 w-full">
<Button
className="hover:bg-[#eef2ff] dark:hover:bg-zinc-800 dark:data-[active]:bg-zinc-800 flex items-center justify-start gap-1 rounded-lg px-2.5 py-2 text-start bg-[#4f46e5] text-white dark:bg-[#6366f1]"
className="hover:bg-zinc-600 w-full dark:hover:bg-zinc-800 dark:data-[active]:bg-zinc-800 flex items-center justify-start gap-1 rounded-lg px-2.5 py-2 text-start bg-[#4f46e5] text-white dark:bg-[#6366f1]"
variant="default"
>
<PlusIcon className="w-4 h-4" />
New Thread
</Button>
<Button
className="hover:bg-zinc-600 w-full dark:hover:bg-zinc-700 dark:data-[active]:bg-zinc-800 flex items-center justify-start gap-1 rounded-lg px-2.5 py-2 text-start bg-zinc-800 text-white"
onClick={toggleDarkMode}
aria-label="Toggle theme"
>
{isDarkMode ? (
<div className="flex items-center gap-2">
<Sun className="w-6 h-6" />
<span>Toggle Light Mode</span>
</div>
) : (
<div className="flex items-center gap-2">
<Moon className="w-6 h-6" />
<span>Toggle Dark Mode</span>
</div>
)}
</Button>
<GithubButton url="https://github.com/mem0ai/mem0/tree/main/examples" className="w-full rounded-lg h-9 pl-2 text-sm font-semibold bg-zinc-800 dark:border-zinc-800 dark:text-white text-white hover:bg-zinc-900" text="View on Github" />
<Link
href={"https://app.mem0.ai/"}
target="_blank"
className="py-2 px-4 w-full rounded-lg h-9 pl-3 text-sm font-semibold dark:bg-zinc-800 dark:hover:bg-zinc-700 bg-zinc-800 text-white hover:bg-zinc-900 dark:text-white"
>
<span className="flex items-center gap-2">
<SaveIcon className="w-4 h-4" />
Save Memories
</span>
</Link>
</div>
</ThreadListPrimitive.New>
<div className="mt-4 mb-2">
<h2 className="text-sm font-medium text-[#475569] dark:text-zinc-300 px-2.5">
@@ -156,29 +192,18 @@ export const Thread: FC<ThreadProps> = ({
</div>
<ThreadListPrimitive.Items components={{ ThreadListItem }} />
</ThreadListPrimitive.Root>
<div>
<Link
href="https://www.assistant-ui.com/"
target="_blank"
className="flex justify-center items-center gap-2"
>
<h1 className="text-sm text-[#475569] dark:text-zinc-300 text-center">
built using
</h1>
<ThemeAwareLogo width={24} height={24} isDarkMode={isDarkMode} />
<p className="text-md font-bold dark:text-zinc-300">
assistant-ui
</p>
</Link>
</div>
</div>
</div>
</div>
</div>
<ScrollArea className="flex-1">
<div className="flex h-full flex-col items-center px-4 pt-8 justify-end">
<ThreadWelcome />
<ScrollArea className="flex-1 w-full">
<div className="flex h-full flex-col w-full items-center px-4 pt-8 justify-end">
<ThreadWelcome
composerInputRef={
composerInputRef as React.RefObject<HTMLTextAreaElement>
}
/>
<ThreadPrimitive.Messages
components={{
@@ -194,9 +219,13 @@ export const Thread: FC<ThreadProps> = ({
</div>
</ScrollArea>
<div className="sticky bottom-0 mt-3 flex w-full max-w-[var(--thread-max-width)] flex-col items-center justify-end rounded-t-lg bg-inherit px-4 pb-4 mx-auto">
<div className="sticky bottom-0 flex w-full max-w-[var(--thread-max-width)] flex-col items-center justify-end rounded-t-lg bg-inherit px-4 md:pb-4 mx-auto">
<ThreadScrollToBottom />
<Composer />
<Composer
composerInputRef={
composerInputRef as React.RefObject<HTMLTextAreaElement>
}
/>
</div>
</ThreadPrimitive.Root>
);
@@ -216,57 +245,74 @@ const ThreadScrollToBottom: FC = () => {
);
};
const ThreadWelcome: FC = () => {
interface ThreadWelcomeProps {
composerInputRef: React.RefObject<HTMLTextAreaElement>;
}
const ThreadWelcome: FC<ThreadWelcomeProps> = ({ composerInputRef }) => {
return (
<ThreadPrimitive.Empty>
<div className="flex w-full max-w-[var(--thread-max-width)] flex-grow flex-col">
<div className="flex w-full flex-grow flex-col items-center justify-start h-[calc(100vh-23rem)] md:h-[calc(100vh-18rem)]">
<div className="flex w-full flex-grow flex-col mt-8 md:h-[calc(100vh-15rem)]">
<div className="flex w-full flex-grow flex-col items-center justify-start">
<div className="flex flex-col items-center justify-center h-full">
<div className="text-2xl md:text-4xl font-bold text-[#1e293b] dark:text-white mb-2">
<div className="text-[2rem] leading-[1] tracking-[-0.02em] md:text-4xl font-bold text-[#1e293b] dark:text-white mb-2 text-center md:w-full w-5/6">
Mem0 - ChatGPT with memory
</div>
<p className="text-center text-sm text-[#1e293b] dark:text-white mb-2 w-3/4">
<p className="text-center text-md text-[#1e293b] dark:text-white mb-2 md:w-3/4 w-5/6">
A personalized AI chat app powered by Mem0 that remembers your
preferences, facts, and memories.
</p>
</div>
</div>
<div className="flex flex-col items-center justify-center">
<div className="flex flex-col items-center justify-center mt-16">
<p className="mt-4 font-medium text-[#1e293b] dark:text-white">
How can I help you today?
</p>
<ThreadWelcomeSuggestions />
<ThreadWelcomeSuggestions composerInputRef={composerInputRef} />
</div>
</div>
</ThreadPrimitive.Empty>
);
};
const ThreadWelcomeSuggestions: FC = () => {
interface ThreadWelcomeSuggestionsProps {
composerInputRef: React.RefObject<HTMLTextAreaElement>;
}
const ThreadWelcomeSuggestions: FC<ThreadWelcomeSuggestionsProps> = ({ composerInputRef }) => {
return (
<div className="mt-3 flex w-full items-stretch justify-center gap-4 dark:text-white">
<div className="mt-3 flex flex-col md:flex-row w-full md:items-stretch justify-center gap-4 dark:text-white items-center">
<ThreadPrimitive.Suggestion
className="hover:bg-[#eef2ff] dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
className="hover:bg-[#eef2ff] w-full dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
prompt="I like to travel to "
method="replace"
onClick={() => {
composerInputRef.current?.focus();
}}
>
<span className="line-clamp-2 text-ellipsis text-sm font-semibold">
Travel
</span>
</ThreadPrimitive.Suggestion>
<ThreadPrimitive.Suggestion
className="hover:bg-[#eef2ff] dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
className="hover:bg-[#eef2ff] w-full dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
prompt="I like to eat "
method="replace"
onClick={() => {
composerInputRef.current?.focus();
}}
>
<span className="line-clamp-2 text-ellipsis text-sm font-semibold">
Food
</span>
</ThreadPrimitive.Suggestion>
<ThreadPrimitive.Suggestion
className="hover:bg-[#eef2ff] dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
className="hover:bg-[#eef2ff] w-full dark:hover:bg-zinc-800 flex max-w-sm grow basis-0 flex-col items-center justify-center rounded-[2rem] border border-[#e2e8f0] dark:border-zinc-700 p-3 transition-colors ease-in"
prompt="I am working on "
method="replace"
onClick={() => {
composerInputRef.current?.focus();
}}
>
<span className="line-clamp-2 text-ellipsis text-sm font-semibold">
Project details
@@ -276,7 +322,11 @@ const ThreadWelcomeSuggestions: FC = () => {
);
};
const Composer: FC = () => {
interface ComposerProps {
composerInputRef: React.RefObject<HTMLTextAreaElement>;
}
const Composer: FC<ComposerProps> = ({ composerInputRef }) => {
return (
<ComposerPrimitive.Root className="focus-within:border-[#4f46e5]/20 dark:focus-within:border-[#6366f1]/20 flex w-full flex-wrap items-end rounded-full border border-[#e2e8f0] dark:border-zinc-700 bg-white dark:bg-zinc-800 px-2.5 shadow-sm transition-colors ease-in">
<ComposerPrimitive.Input
@@ -284,6 +334,7 @@ const Composer: FC = () => {
autoFocus
placeholder="Message to Mem0..."
className="placeholder:text-zinc-400 dark:placeholder:text-zinc-500 max-h-40 flex-grow resize-none border-none bg-transparent px-2 py-4 text-sm outline-none focus:ring-0 disabled:cursor-not-allowed text-[#1e293b] dark:text-zinc-200"
ref={composerInputRef}
/>
<ComposerAction />
</ComposerPrimitive.Root>
@@ -1,18 +1,18 @@
import { cn } from "@/lib/utils";
const GithubButton = ({ url }: { url: string }) => {
const GithubButton = ({ url, className, text }: { url: string, className?: string, text?: string }) => {
return (
<a
href={url}
target="_blank"
rel="noopener noreferrer"
className="flex items-center bg-black text-white rounded-full shadow-lg hover:bg-gray-800 transition border border-gray-700"
className={cn("flex items-center bg-black text-white rounded-full shadow-lg hover:bg-gray-800 transition border border-gray-700", className)}
>
<svg
xmlns="http://www.w3.org/2000/svg"
viewBox="0 0 24 24"
fill="white"
className="w-6 h-6"
className="w-5 h-5 md:w-6 md:h-6"
>
<path
fillRule="evenodd"
@@ -20,6 +20,7 @@ const GithubButton = ({ url }: { url: string }) => {
clipRule="evenodd"
/>
</svg>
{text && <span className="ml-2">{text}</span>}
</a>
);
};
+83
View File
@@ -0,0 +1,83 @@
import MemoryClient from "mem0ai";
import { OpenAI } from "openai";
import { zodResponsesFunction } from "openai/helpers/zod";
import { z } from "zod";
const mem0Config = {
apiKey: process.env.MEM0_API_KEY, // GET THIS API KEY FROM MEM0 (https://app.mem0.ai/dashboard/api-keys)
user_id: "sample-user",
};
async function run() {
// RESPONES WITHOUT MEMORIES
console.log("\n\nRESPONES WITHOUT MEMORIES\n\n");
await main();
// ADDING SOME SAMPLE MEMORIES
await addSampleMemories();
// RESPONES WITH MEMORIES
console.log("\n\nRESPONES WITH MEMORIES\n\n");
await main(true);
}
// OpenAI Response Schema
const CarSchema = z.object({
car_name: z.string(),
car_price: z.string(),
car_url: z.string(),
car_image: z.string(),
car_description: z.string(),
});
const Cars = z.object({
cars: z.array(CarSchema),
});
async function main(memory = false) {
const openAIClient = new OpenAI();
const mem0Client = new MemoryClient(mem0Config);
const input = "Suggest me some cars that I can buy today.";
const tool = zodResponsesFunction({ name: "carRecommendations", parameters: Cars });
// First, let's store the user's memories from user input if any
await mem0Client.add([{
role: "user",
content: input,
}], mem0Config);
// Then search for relevant memories
let relevantMemories = []
if (memory) {
relevantMemories = await mem0Client.search(input, mem0Config);
}
const response = await openAIClient.responses.create({
model: "gpt-4o",
tools: [{ type: "web_search_preview" }, tool],
input: `${getMemoryString(relevantMemories)}\n${input}`,
});
console.log(response.output);
}
async function addSampleMemories() {
const mem0Client = new MemoryClient(mem0Config);
const myInterests = "I Love BMW, Audi and Porsche. I Hate Mercedes. I love Red cars and Maroon cars. I have a budget of 120K to 150K USD. I like Audi the most.";
await mem0Client.add([{
role: "user",
content: myInterests,
}], mem0Config);
}
const getMemoryString = (memories) => {
const MEMORY_STRING_PREFIX = "These are the memories I have stored. Give more weightage to the question by users and try to answer that first. You have to modify your answer based on the memories I have provided. If the memories are irrelevant you can ignore them. Also don't reply to this section of the prompt, or the memories, they are only for your reference. The MEMORIES of the USER are: \n\n";
const memoryString = memories.map((mem) => `${mem.memory}`).join("\n") ?? "";
return memoryString.length > 0 ? `${MEMORY_STRING_PREFIX}${memoryString}` : "";
};
run().catch(console.error);
@@ -0,0 +1,19 @@
{
"name": "openai-inbuilt-tools",
"version": "1.0.0",
"description": "",
"license": "ISC",
"author": "",
"type": "module",
"main": "index.js",
"scripts": {
"test": "echo \"Error: no test specified\" && exit 1",
"start": "node index.js"
},
"packageManager": "pnpm@10.5.2+sha512.da9dc28cd3ff40d0592188235ab25d3202add8a207afbedc682220e4a0029ffbff4562102b9e6e46b4e3f9e8bd53e6d05de48544b0c57d4b0179e22c76d1199b",
"dependencies": {
"mem0ai": "^2.1.2",
"openai": "^4.87.2",
"zod": "^3.24.2"
}
}
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "mem0ai",
"version": "2.1.1",
"version": "2.1.4",
"description": "The Memory Layer For Your AI Apps",
"main": "./dist/index.js",
"module": "./dist/index.mjs",
+4
View File
@@ -205,6 +205,10 @@ export default class MemoryClient {
if (options.project_name) delete options.project_name;
}
if (options.api_version) {
options.version = options.api_version.toString();
}
const payload = this._preparePayload(messages, options);
const response = await this._fetchWithErrorHandling(
`${this.host}/v1/memories/`,
+3
View File
@@ -1,4 +1,6 @@
export interface MemoryOptions {
api_version?: API_VERSION | string;
version?: API_VERSION | string;
user_id?: string;
agent_id?: string;
app_id?: string;
@@ -17,6 +19,7 @@ export interface MemoryOptions {
enable_graph?: boolean;
start_date?: string;
end_date?: string;
custom_categories?: custom_categories[];
}
export interface ProjectOptions {
+1 -1
View File
@@ -4,7 +4,7 @@ import type { TelemetryClient } from "./telemetry.types";
let version = "1.0.20";
const MEM0_TELEMETRY = process.env.MEM0_TELEMETRY !== "false";
const MEM0_TELEMETRY = "false";
const POSTHOG_API_KEY = "phc_hgJkUVJFYtmaJqrvf6CYN67TIQ8yhXAkWzUn9AMU4yX";
const POSTHOG_HOST = "https://us.i.posthog.com";
+2 -2
View File
@@ -6,7 +6,7 @@ from typing import Any, Dict, List, Optional, Union
import httpx
from mem0.memory.setup import setup_config, get_user_id
from mem0.memory.setup import get_user_id, setup_config
from mem0.memory.telemetry import capture_client_event
logger = logging.getLogger(__name__)
@@ -110,7 +110,7 @@ class MemoryClient:
try:
error_data = e.response.json()
error_message = error_data.get("detail", str(e))
except:
except Exception:
error_message = str(e)
raise ValueError(f"Error: {error_message}")
+35 -9
View File
@@ -1,27 +1,53 @@
from typing import Any, Dict
from typing import Any, Dict, Optional
from pydantic import BaseModel, Field, model_validator
class AzureAISearchConfig(BaseModel):
collection_name: str = Field("mem0", description="Name of the collection")
service_name: str = Field(None, description="Azure Cognitive Search service name")
api_key: str = Field(None, description="API key for the Azure Cognitive Search service")
service_name: str = Field(None, description="Azure AI Search service name")
api_key: str = Field(None, description="API key for the Azure AI Search service")
embedding_model_dims: int = Field(None, description="Dimension of the embedding vector")
use_compression: bool = Field(False, description="Whether to use scalar quantization vector compression.")
compression_type: Optional[str] = Field(
None,
description="Type of vector compression to use. Options: 'scalar', 'binary', or None"
)
use_float16: bool = Field(
False,
description="Whether to store vectors in half precision (Edm.Half) instead of full precision (Edm.Single)"
)
@model_validator(mode="before")
@classmethod
def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
allowed_fields = set(cls.model_fields.keys())
input_fields = set(values.keys())
extra_fields = input_fields - allowed_fields
# Check for use_compression to provide a helpful error
if "use_compression" in extra_fields:
raise ValueError(
"The parameter 'use_compression' is no longer supported. "
"Please use 'compression_type=\"scalar\"' instead of 'use_compression=True' "
"or 'compression_type=None' instead of 'use_compression=False'."
)
if extra_fields:
raise ValueError(
f"Extra fields not allowed: {', '.join(extra_fields)}. Please input only the following fields: {', '.join(allowed_fields)}"
f"Extra fields not allowed: {', '.join(extra_fields)}. "
f"Please input only the following fields: {', '.join(allowed_fields)}"
)
# Validate compression_type values
if "compression_type" in values and values["compression_type"] is not None:
valid_types = ["scalar", "binary"]
if values["compression_type"].lower() not in valid_types:
raise ValueError(
f"Invalid compression_type: {values['compression_type']}. "
f"Must be one of: {', '.join(valid_types)}, or None"
)
return values
model_config = {
"arbitrary_types_allowed": True,
}
}
+42
View File
@@ -0,0 +1,42 @@
from typing import Any, ClassVar, Dict, Optional
from pydantic import BaseModel, Field, model_validator
class WeaviateConfig(BaseModel):
from weaviate import WeaviateClient
WeaviateClient: ClassVar[type] = WeaviateClient
collection_name: str = Field("mem0", description="Name of the collection")
embedding_model_dims: int = Field(1536, description="Dimensions of the embedding model")
cluster_url: Optional[str] = Field(None, description="URL for Weaviate server")
auth_client_secret: Optional[str] = Field(None, description="API key for Weaviate authentication")
additional_headers: Optional[Dict[str, str]] = Field(None, description="Additional headers for requests")
@model_validator(mode="before")
@classmethod
def check_connection_params(cls, values: Dict[str, Any]) -> Dict[str, Any]:
cluster_url = values.get("cluster_url")
if not cluster_url:
raise ValueError("'cluster_url' must be provided.")
return values
@model_validator(mode="before")
@classmethod
def validate_extra_fields(cls, values: Dict[str, Any]) -> Dict[str, Any]:
allowed_fields = set(cls.model_fields.keys())
input_fields = set(values.keys())
extra_fields = input_fields - allowed_fields
if extra_fields:
raise ValueError(
f"Extra fields not allowed: {', '.join(extra_fields)}. Please input only the following fields: {', '.join(allowed_fields)}"
)
return values
model_config = {
"arbitrary_types_allowed": True,
}
+18 -15
View File
@@ -4,14 +4,26 @@ from typing import Dict, List, Optional
try:
import anthropic
except ImportError:
raise ImportError("The 'anthropic' library is required. Please install it using 'pip install anthropic'.")
raise ImportError(
"The 'anthropic' library is required. Please install it using 'pip install anthropic'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class AnthropicLLM(LLMBase):
"""
A class for interacting with Anthropic's Claude models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the AnthropicLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
@@ -23,23 +35,17 @@ class AnthropicLLM(LLMBase):
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
) -> str:
"""
Generate a response based on the given messages using Anthropic.
Generates a response using Anthropic's Claude model based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
Returns:
str: The generated response.
str: The generated response from the model.
"""
# Separate system message from other messages
# Extract system message separately
system_message = ""
filtered_messages = []
for message in messages:
@@ -56,9 +62,6 @@ class AnthropicLLM(LLMBase):
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if tools: # TODO: Remove tools if no issues found with new memory addition logic
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.messages.create(**params)
return response.content[0].text
+50 -126
View File
@@ -4,14 +4,26 @@ from typing import Any, Dict, List, Optional
try:
import boto3
except ImportError:
raise ImportError("The 'boto3' library is required. Please install it using 'pip install boto3'.")
raise ImportError(
"The 'boto3' library is required. Please install it using 'pip install boto3'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class AWSBedrockLLM(LLMBase):
"""
A wrapper for AWS Bedrock's language models, integrating them with the LLMBase class.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the AWS Bedrock LLM with the provided configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration object for the model.
"""
super().__init__(config)
if not self.config.model:
@@ -25,49 +37,29 @@ class AWSBedrockLLM(LLMBase):
def _format_messages(self, messages: List[Dict[str, str]]) -> str:
"""
Formats a list of messages into the required prompt structure for the model.
Formats a list of messages into a structured prompt for the model.
Args:
messages (List[Dict[str, str]]): A list of dictionaries where each dictionary represents a message.
Each dictionary contains 'role' and 'content' keys.
messages (List[Dict[str, str]]): A list of dictionaries containing 'role' and 'content'.
Returns:
str: A formatted string combining all messages, structured with roles capitalized and separated by newlines.
"""
formatted_messages = []
for message in messages:
role = message["role"].capitalize()
content = message["content"]
formatted_messages.append(f"\n\n{role}: {content}")
formatted_messages = [
f"\n\n{msg['role'].capitalize()}: {msg['content']}" for msg in messages
]
return "".join(formatted_messages) + "\n\nAssistant:"
def _parse_response(self, response, tools) -> str:
def _parse_response(self, response) -> str:
"""
Process the response based on whether tools are used or not.
Extracts the generated response from the API response.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
response: The raw response from the AWS Bedrock API.
Returns:
str or dict: The processed response.
str: The generated response text.
"""
if tools:
processed_response = {"tool_calls": []}
if response["output"]["message"]["content"]:
for item in response["output"]["message"]["content"]:
if "toolUse" in item:
processed_response["tool_calls"].append(
{
"name": item["toolUse"]["name"],
"arguments": item["toolUse"]["input"],
}
)
return processed_response
response_body = json.loads(response["body"].read().decode())
return response_body.get("completion", "")
@@ -76,22 +68,21 @@ class AWSBedrockLLM(LLMBase):
provider: str,
model: str,
prompt: str,
model_kwargs: Optional[Dict[str, Any]] = {},
model_kwargs: Optional[Dict[str, Any]] = None,
) -> Dict[str, Any]:
"""
Prepares the input dictionary for the specified provider's model by mapping and renaming
keys in the input based on the provider's requirements.
Prepares the input dictionary for the specified provider's model.
Args:
provider (str): The name of the service provider (e.g., "meta", "ai21", "mistral", "cohere", "amazon").
model (str): The name or identifier of the model being used.
prompt (str): The text prompt to be processed by the model.
model_kwargs (Dict[str, Any]): Additional keyword arguments specific to the model's requirements.
provider (str): The model provider (e.g., "meta", "ai21", "mistral", "cohere", "amazon").
model (str): The model identifier.
prompt (str): The input prompt.
model_kwargs (Optional[Dict[str, Any]]): Additional model parameters.
Returns:
Dict[str, Any]: The prepared input dictionary with the correct keys and values for the specified provider.
Dict[str, Any]: The prepared input dictionary.
"""
model_kwargs = model_kwargs or {}
input_body = {"prompt": prompt, **model_kwargs}
provider_mappings = {
@@ -119,102 +110,35 @@ class AWSBedrockLLM(LLMBase):
},
}
input_body["textGenerationConfig"] = {
k: v for k, v in input_body["textGenerationConfig"].items() if v is not None
k: v
for k, v in input_body["textGenerationConfig"].items()
if v is not None
}
return input_body
def _convert_tool_format(self, original_tools):
def generate_response(self, messages: List[Dict[str, str]]) -> str:
"""
Converts a list of tools from their original format to a new standardized format.
Generates a response using AWS Bedrock based on the provided messages.
Args:
original_tools (list): A list of dictionaries representing the original tools, each containing a 'type' key and corresponding details.
messages (List[Dict[str, str]]): List of message dictionaries containing 'role' and 'content'.
Returns:
list: A list of dictionaries representing the tools in the new standardized format.
str: The generated response text.
"""
new_tools = []
prompt = self._format_messages(messages)
provider = self.config.model.split(".")[0]
input_body = self._prepare_input(
provider, self.config.model, prompt, self.model_kwargs
)
body = json.dumps(input_body)
for tool in original_tools:
if tool["type"] == "function":
function = tool["function"]
new_tool = {
"toolSpec": {
"name": function["name"],
"description": function["description"],
"inputSchema": {
"json": {
"type": "object",
"properties": {},
"required": function["parameters"].get("required", []),
}
},
}
}
response = self.client.invoke_model(
body=body,
modelId=self.config.model,
accept="application/json",
contentType="application/json",
)
for prop, details in function["parameters"].get("properties", {}).items():
new_tool["toolSpec"]["inputSchema"]["json"]["properties"][prop] = {
"type": details.get("type", "string"),
"description": details.get("description", ""),
}
new_tools.append(new_tool)
return new_tools
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generate a response based on the given messages using AWS Bedrock.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
Returns:
str: The generated response.
"""
if tools:
# Use converse method when tools are provided
messages = [
{
"role": "user",
"content": [{"text": message["content"]} for message in messages],
}
]
inference_config = {
"temperature": self.model_kwargs["temperature"],
"maxTokens": self.model_kwargs["max_tokens_to_sample"],
"topP": self.model_kwargs["top_p"],
}
tools_config = {"tools": self._convert_tool_format(tools)}
response = self.client.converse(
modelId=self.config.model,
messages=messages,
inferenceConfig=inference_config,
toolConfig=tools_config,
)
else:
# Use invoke_model method when no tools are provided
prompt = self._format_messages(messages)
provider = self.model.split(".")[0]
input_body = self._prepare_input(provider, self.config.model, prompt, **self.model_kwargs)
body = json.dumps(input_body)
response = self.client.invoke_model(
body=body,
modelId=self.model,
accept="application/json",
contentType="application/json",
)
return self._parse_response(response, tools)
return self._parse_response(response)
+31 -50
View File
@@ -9,17 +9,35 @@ from mem0.llms.base import LLMBase
class AzureOpenAILLM(LLMBase):
"""
A class for interacting with Azure OpenAI models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the AzureOpenAILLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
# Model name should match the custom deployment name chosen for it.
# Ensure model name is set; it should match the Azure OpenAI deployment name.
if not self.config.model:
self.config.model = "gpt-4o"
api_key = self.config.azure_kwargs.api_key or os.getenv("LLM_AZURE_OPENAI_API_KEY")
azure_deployment = self.config.azure_kwargs.azure_deployment or os.getenv("LLM_AZURE_DEPLOYMENT")
azure_endpoint = self.config.azure_kwargs.azure_endpoint or os.getenv("LLM_AZURE_ENDPOINT")
api_version = self.config.azure_kwargs.api_version or os.getenv("LLM_AZURE_API_VERSION")
api_key = self.config.azure_kwargs.api_key or os.getenv(
"LLM_AZURE_OPENAI_API_KEY"
)
azure_deployment = self.config.azure_kwargs.azure_deployment or os.getenv(
"LLM_AZURE_DEPLOYMENT"
)
azure_endpoint = self.config.azure_kwargs.azure_endpoint or os.getenv(
"LLM_AZURE_ENDPOINT"
)
api_version = self.config.azure_kwargs.api_version or os.getenv(
"LLM_AZURE_API_VERSION"
)
default_headers = self.config.azure_kwargs.default_headers
self.client = AzureOpenAI(
@@ -31,54 +49,20 @@ class AzureOpenAILLM(LLMBase):
default_headers=default_headers,
)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using Azure OpenAI.
Generates a response using Azure OpenAI based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -87,11 +71,8 @@ class AzureOpenAILLM(LLMBase):
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+32 -50
View File
@@ -9,20 +9,38 @@ from mem0.llms.base import LLMBase
class AzureOpenAIStructuredLLM(LLMBase):
"""
A class for interacting with Azure OpenAI models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the AzureOpenAIStructuredLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
# Model name should match the custom deployment name chosen for it.
# Ensure model name is set; it should match the Azure OpenAI deployment name.
if not self.config.model:
self.config.model = "gpt-4o-2024-08-06"
api_key = os.getenv("LLM_AZURE_OPENAI_API_KEY") or self.config.azure_kwargs.api_key
azure_deployment = os.getenv("LLM_AZURE_DEPLOYMENT") or self.config.azure_kwargs.azure_deployment
azure_endpoint = os.getenv("LLM_AZURE_ENDPOINT") or self.config.azure_kwargs.azure_endpoint
api_version = os.getenv("LLM_AZURE_API_VERSION") or self.config.azure_kwargs.api_version
api_key = (
os.getenv("LLM_AZURE_OPENAI_API_KEY") or self.config.azure_kwargs.api_key
)
azure_deployment = (
os.getenv("LLM_AZURE_DEPLOYMENT")
or self.config.azure_kwargs.azure_deployment
)
azure_endpoint = (
os.getenv("LLM_AZURE_ENDPOINT") or self.config.azure_kwargs.azure_endpoint
)
api_version = (
os.getenv("LLM_AZURE_API_VERSION") or self.config.azure_kwargs.api_version
)
default_headers = self.config.azure_kwargs.default_headers
# Can display a warning if API version is of model and api-version
self.client = AzureOpenAI(
azure_deployment=azure_deployment,
azure_endpoint=azure_endpoint,
@@ -32,54 +50,20 @@ class AzureOpenAIStructuredLLM(LLMBase):
default_headers=default_headers,
)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using Azure OpenAI.
Generates a response using Azure OpenAI based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -88,11 +72,9 @@ class AzureOpenAIStructuredLLM(LLMBase):
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if response_format:
params["response_format"] = response_format
if tools:
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+6 -2
View File
@@ -4,8 +4,12 @@ from pydantic import BaseModel, Field, field_validator
class LlmConfig(BaseModel):
provider: str = Field(description="Provider of the LLM (e.g., 'ollama', 'openai')", default="openai")
config: Optional[dict] = Field(description="Configuration for the specific LLM", default={})
provider: str = Field(
description="Provider of the LLM (e.g., 'ollama', 'openai')", default="openai"
)
config: Optional[dict] = Field(
description="Configuration for the specific LLM", default={}
)
@field_validator("config")
def validate_config(cls, v, values):
+20 -47
View File
@@ -9,64 +9,42 @@ from mem0.llms.base import LLMBase
class DeepSeekLLM(LLMBase):
"""
A class for interacting with DeepSeek's language models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the DeepSeekLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
self.config.model = "deepseek-chat"
api_key = self.config.api_key or os.getenv("DEEPSEEK_API_KEY")
base_url = self.config.deepseek_base_url or os.getenv("DEEPSEEK_API_BASE") or "https://api.deepseek.com"
base_url = (
self.config.deepseek_base_url
or os.getenv("DEEPSEEK_API_BASE")
or "https://api.deepseek.com"
)
self.client = OpenAI(api_key=api_key, base_url=base_url)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
) -> str:
"""
Generate a response based on the given messages using DeepSeek.
Generates a response using DeepSeek based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -75,10 +53,5 @@ class DeepSeekLLM(LLMBase):
"max_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if tools:
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+28 -102
View File
@@ -3,8 +3,7 @@ from typing import Dict, List, Optional
try:
import google.generativeai as genai
from google.generativeai import GenerativeModel, protos
from google.generativeai.types import content_types
from google.generativeai import GenerativeModel
except ImportError:
raise ImportError(
"The 'google-generativeai' library is required. Please install it using 'pip install google-generativeai'."
@@ -15,7 +14,17 @@ from mem0.llms.base import LLMBase
class GeminiLLM(LLMBase):
"""
A wrapper for Google's Gemini language model, integrating it with the LLMBase class.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the Gemini LLM with the provided configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration object for the model.
"""
super().__init__(config)
if not self.config.model:
@@ -25,51 +34,25 @@ class GeminiLLM(LLMBase):
genai.configure(api_key=api_key)
self.client = GenerativeModel(model_name=self.config.model)
def _parse_response(self, response, tools):
def _reformat_messages(
self, messages: List[Dict[str, str]]
) -> List[Dict[str, str]]:
"""
Process the response based on whether tools are used or not.
Reformats messages to match the Gemini API's expected structure.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
messages (List[Dict[str, str]]): A list of messages with 'role' and 'content' keys.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": (content if (content := response.candidates[0].content.parts[0].text) else None),
"tool_calls": [],
}
for part in response.candidates[0].content.parts:
if fn := part.function_call:
if isinstance(fn, protos.FunctionCall):
fn_call = type(fn).to_dict(fn)
processed_response["tool_calls"].append({"name": fn_call["name"], "arguments": fn_call["args"]})
continue
processed_response["tool_calls"].append({"name": fn.name, "arguments": fn.args})
return processed_response
else:
return response.candidates[0].content.parts[0].text
def _reformat_messages(self, messages: List[Dict[str, str]]):
"""
Reformat messages for Gemini.
Args:
messages: The list of messages provided in the request.
Returns:
list: The list of messages in the required format.
List[Dict[str, str]]: Reformatted messages in the required format.
"""
new_messages = []
for message in messages:
if message["role"] == "system":
content = "THIS IS A SYSTEM PROMPT. YOU MUST OBEY THIS: " + message["content"]
content = (
"THIS IS A SYSTEM PROMPT. YOU MUST OBEY THIS: " + message["content"]
)
else:
content = message["content"]
@@ -82,90 +65,33 @@ class GeminiLLM(LLMBase):
return new_messages
def _reformat_tools(self, tools: Optional[List[Dict]]):
"""
Reformat tools for Gemini.
Args:
tools: The list of tools provided in the request.
Returns:
list: The list of tools in the required format.
"""
def remove_additional_properties(data):
"""Recursively removes 'additionalProperties' from nested dictionaries."""
if isinstance(data, dict):
filtered_dict = {
key: remove_additional_properties(value)
for key, value in data.items()
if not (key == "additionalProperties")
}
return filtered_dict
else:
return data
new_tools = []
if tools:
for tool in tools:
func = tool["function"].copy()
new_tools.append({"function_declarations": [remove_additional_properties(func)]})
# TODO: temporarily ignore it to pass tests, will come back to update according to standards later.
# return content_types.to_function_library(new_tools)
return new_tools
else:
return None
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
self, messages: List[Dict[str, str]], response_format: Optional[Dict] = None
) -> str:
"""
Generate a response based on the given messages using Gemini.
Generates a response from Gemini based on the given conversation history.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format for the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): List of message dictionaries containing 'role' and 'content'.
response_format (Optional[Dict]): Specifies the response format (e.g., JSON schema).
Returns:
str: The generated response.
str: The generated response as text.
"""
params = {
"temperature": self.config.temperature,
"max_output_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if response_format is not None and response_format["type"] == "json_object":
if response_format and response_format.get("type") == "json_object":
params["response_mime_type"] = "application/json"
if "schema" in response_format:
params["response_schema"] = response_format["schema"]
if tool_choice:
tool_config = content_types.to_tool_config(
{
"function_calling_config": {
"mode": tool_choice,
"allowed_function_names": (
[tool["function"]["name"] for tool in tools] if tool_choice == "any" else None
),
}
}
)
response = self.client.generate_content(
contents=self._reformat_messages(messages),
tools=self._reformat_tools(tools),
generation_config=genai.GenerationConfig(**params),
tool_config=tool_config,
)
return self._parse_response(response, tools)
return response.candidates[0].content.parts[0].text
+20 -46
View File
@@ -5,14 +5,26 @@ from typing import Dict, List, Optional
try:
from groq import Groq
except ImportError:
raise ImportError("The 'groq' library is required. Please install it using 'pip install groq'.")
raise ImportError(
"The 'groq' library is required. Please install it using 'pip install groq'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class GroqLLM(LLMBase):
"""
A class for interacting with Groq's language models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the GroqLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
@@ -21,54 +33,20 @@ class GroqLLM(LLMBase):
api_key = self.config.api_key or os.getenv("GROQ_API_KEY")
self.client = Groq(api_key=api_key)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using Groq.
Generates a response using Groq based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -79,9 +57,5 @@ class GroqLLM(LLMBase):
}
if response_format:
params["response_format"] = response_format
if tools:
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+23 -46
View File
@@ -4,70 +4,50 @@ from typing import Dict, List, Optional
try:
import litellm
except ImportError:
raise ImportError("The 'litellm' library is required. Please install it using 'pip install litellm'.")
raise ImportError(
"The 'litellm' library is required. Please install it using 'pip install litellm'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class LiteLLM(LLMBase):
"""
A class for interacting with LiteLLM's language models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the LiteLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
self.config.model = "gpt-4o-mini"
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using Litellm.
Generates a response using LiteLLM based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
if not litellm.supports_function_calling(self.config.model):
raise ValueError(f"Model '{self.config.model}' in litellm does not support function calling.")
raise ValueError(
f"Model '{self.config.model}' in LiteLLM does not support function calling."
)
params = {
"model": self.config.model,
@@ -78,9 +58,6 @@ class LiteLLM(LLMBase):
}
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
params["tools"] = tools
params["tool_choice"] = tool_choice
response = litellm.completion(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+22 -46
View File
@@ -3,77 +3,56 @@ from typing import Dict, List, Optional
try:
from ollama import Client
except ImportError:
raise ImportError("The 'ollama' library is required. Please install it using 'pip install ollama'.")
raise ImportError(
"The 'ollama' library is required. Please install it using 'pip install ollama'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class OllamaLLM(LLMBase):
"""
A class for interacting with Ollama's language models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the OllamaLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
self.config.model = "llama3.1:70b"
self.client = Client(host=self.config.ollama_base_url)
self._ensure_model_exists()
def _ensure_model_exists(self):
"""
Ensure the specified model exists locally. If not, pull it from Ollama.
Ensures the specified model exists locally. If not, pulls it from Ollama.
"""
local_models = self.client.list()["models"]
if not any(model.get("name") == self.config.model for model in local_models):
self.client.pull(self.config.model)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response["message"]["content"],
"tool_calls": [],
}
if response["message"].get("tool_calls"):
for tool_call in response["message"]["tool_calls"]:
processed_response["tool_calls"].append(
{
"name": tool_call["function"]["name"],
"arguments": tool_call["function"]["arguments"],
}
)
return processed_response
else:
return response["message"]["content"]
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using OpenAI.
Generates a response using Ollama based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -87,8 +66,5 @@ class OllamaLLM(LLMBase):
if response_format:
params["format"] = "json"
if tools:
params["tools"] = tools
response = self.client.chat(**params)
return self._parse_response(response, tools)
return response["message"]["content"]
+22 -45
View File
@@ -9,7 +9,17 @@ from mem0.llms.base import LLMBase
class OpenAILLM(LLMBase):
"""
A class to interact with OpenAI or OpenRouter APIs for generating responses using LLMs.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the OpenAILLM instance.
Args:
config (Optional[BaseLlmConfig]): Configuration for the LLM, including model, API key, and base URLs.
"""
super().__init__(config)
if not self.config.model:
@@ -24,57 +34,27 @@ class OpenAILLM(LLMBase):
)
else:
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
base_url = self.config.openai_base_url or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1"
base_url = (
self.config.openai_base_url
or os.getenv("OPENAI_API_BASE")
or "https://api.openai.com/v1"
)
self.client = OpenAI(api_key=api_key, base_url=base_url)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using OpenAI.
Generates a response based on the provided messages using OpenAI or OpenRouter.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of message dictionaries containing 'role' and 'content'.
response_format (Optional[str]): The format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -102,9 +82,6 @@ class OpenAILLM(LLMBase):
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+23 -48
View File
@@ -9,66 +9,45 @@ from mem0.llms.base import LLMBase
class OpenAIStructuredLLM(LLMBase):
"""
A class for interacting with OpenAI's structured language models using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the OpenAIStructuredLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
self.config.model = "gpt-4o-2024-08-06"
api_key = self.config.api_key or os.getenv("OPENAI_API_KEY")
base_url = self.config.openai_base_url or os.getenv("OPENAI_API_BASE") or "https://api.openai.com/v1"
base_url = (
self.config.openai_base_url
or os.getenv("OPENAI_API_BASE")
or "https://api.openai.com/v1"
)
self.client = OpenAI(api_key=api_key, base_url=base_url)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools (list, optional): List of tools that the model can call.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using OpenAI.
Generates a response using OpenAI based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -78,10 +57,6 @@ class OpenAIStructuredLLM(LLMBase):
if response_format:
params["response_format"] = response_format
if tools:
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.beta.chat.completions.parse(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+20 -45
View File
@@ -5,14 +5,26 @@ from typing import Dict, List, Optional
try:
from together import Together
except ImportError:
raise ImportError("The 'together' library is required. Please install it using 'pip install together'.")
raise ImportError(
"The 'together' library is required. Please install it using 'pip install together'."
)
from mem0.configs.llms.base import BaseLlmConfig
from mem0.llms.base import LLMBase
class TogetherLLM(LLMBase):
"""
A class for interacting with the TogetherAI language model using the specified configuration.
"""
def __init__(self, config: Optional[BaseLlmConfig] = None):
"""
Initializes the TogetherLLM instance with the given configuration.
Args:
config (Optional[BaseLlmConfig]): Configuration settings for the language model.
"""
super().__init__(config)
if not self.config.model:
@@ -21,54 +33,20 @@ class TogetherLLM(LLMBase):
api_key = self.config.api_key or os.getenv("TOGETHER_API_KEY")
self.client = Together(api_key=api_key)
def _parse_response(self, response, tools):
"""
Process the response based on whether tools are used or not.
Args:
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str or dict: The processed response.
"""
if tools:
processed_response = {
"content": response.choices[0].message.content,
"tool_calls": [],
}
if response.choices[0].message.tool_calls:
for tool_call in response.choices[0].message.tool_calls:
processed_response["tool_calls"].append(
{
"name": tool_call.function.name,
"arguments": json.loads(tool_call.function.arguments),
}
)
return processed_response
else:
return response.choices[0].message.content
def generate_response(
self,
messages: List[Dict[str, str]],
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
response_format: Optional[str] = None,
) -> str:
"""
Generate a response based on the given messages using TogetherAI.
Generates a response using TogetherAI based on the provided messages.
Args:
messages (list): List of message dicts containing 'role' and 'content'.
response_format (str or object, optional): Format of the response. Defaults to "text".
tools (list, optional): List of tools that the model can call. Defaults to None.
tool_choice (str, optional): Tool choice method. Defaults to "auto".
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
response_format (Optional[str]): The desired format of the response. Defaults to None.
Returns:
str: The generated response.
str: The generated response from the model.
"""
params = {
"model": self.config.model,
@@ -79,9 +57,6 @@ class TogetherLLM(LLMBase):
}
if response_format:
params["response_format"] = response_format
if tools: # TODO: Remove tools if no issues found with new memory addition logic
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return self._parse_response(response, tools)
return response.choices[0].message.content
+5 -1
View File
@@ -15,7 +15,11 @@ class XAILLM(LLMBase):
self.config.model = "grok-2-latest"
api_key = self.config.api_key or os.getenv("XAI_API_KEY")
base_url = self.config.xai_base_url or os.getenv("XAI_API_BASE") or "https://api.x.ai/v1"
base_url = (
self.config.xai_base_url
or os.getenv("XAI_API_BASE")
or "https://api.x.ai/v1"
)
self.client = OpenAI(api_key=api_key, base_url=base_url)
def generate_response(self, messages: List[Dict[str, str]], response_format=None):
+7 -10
View File
@@ -3,23 +3,20 @@ import logging
from mem0.memory.utils import format_entities
try:
from langchain_community.graphs import Neo4jGraph
from langchain_neo4j import Neo4jGraph
except ImportError:
raise ImportError("langchain_community is not installed. Please install it using pip install langchain-community")
raise ImportError("langchain_neo4j is not installed. Please install it using pip install langchain-neo4j")
try:
from rank_bm25 import BM25Okapi
except ImportError:
raise ImportError("rank_bm25 is not installed. Please install it using pip install rank-bm25")
from mem0.graphs.tools import (
DELETE_MEMORY_STRUCT_TOOL_GRAPH,
DELETE_MEMORY_TOOL_GRAPH,
EXTRACT_ENTITIES_STRUCT_TOOL,
EXTRACT_ENTITIES_TOOL,
RELATIONS_STRUCT_TOOL,
RELATIONS_TOOL,
)
from mem0.graphs.tools import (DELETE_MEMORY_STRUCT_TOOL_GRAPH,
DELETE_MEMORY_TOOL_GRAPH,
EXTRACT_ENTITIES_STRUCT_TOOL,
EXTRACT_ENTITIES_TOOL, RELATIONS_STRUCT_TOOL,
RELATIONS_TOOL)
from mem0.graphs.utils import EXTRACT_RELATIONS_PROMPT, get_delete_messages
from mem0.utils.factory import EmbedderFactory, LlmFactory
+32 -12
View File
@@ -57,12 +57,27 @@ class Memory(MemoryBase):
@classmethod
def from_config(cls, config_dict: Dict[str, Any]):
try:
config = cls._process_config(config_dict)
config = MemoryConfig(**config_dict)
except ValidationError as e:
logger.error(f"Configuration validation error: {e}")
raise
return cls(config)
@staticmethod
def _process_config(config_dict: Dict[str, Any]) -> Dict[str, Any]:
if "graph_store" in config_dict:
if "vector_store" not in config_dict and "embedder" in config_dict:
config_dict["vector_store"] = {}
config_dict["vector_store"]["config"] = {}
config_dict["vector_store"]["config"]["embedding_model_dims"] = config_dict["embedder"]["config"]["embedding_dims"]
try:
return config_dict
except ValidationError as e:
logger.error(f"Configuration validation error: {e}")
raise
def add(
self,
messages,
@@ -71,6 +86,7 @@ class Memory(MemoryBase):
run_id=None,
metadata=None,
filters=None,
infer=True,
prompt=None,
):
"""
@@ -83,6 +99,7 @@ class Memory(MemoryBase):
run_id (str, optional): ID of the run creating the memory. Defaults to None.
metadata (dict, optional): Metadata to store with the memory. Defaults to None.
filters (dict, optional): Filters to apply to the search. Defaults to None.
infer (bool, optional): Whether to infer the memories. Defaults to True.
prompt (str, optional): Prompt to use for memory deduction. Defaults to None.
Returns:
@@ -121,7 +138,7 @@ class Memory(MemoryBase):
messages = parse_vision_messages(messages)
with concurrent.futures.ThreadPoolExecutor() as executor:
future1 = executor.submit(self._add_to_vector_store, messages, metadata, filters)
future1 = executor.submit(self._add_to_vector_store, messages, metadata, filters, infer)
future2 = executor.submit(self._add_to_graph, messages, filters)
concurrent.futures.wait([future1, future2])
@@ -147,7 +164,16 @@ class Memory(MemoryBase):
return {"results": vector_store_result}
def _add_to_vector_store(self, messages, metadata, filters):
def _add_to_vector_store(self, messages, metadata, filters, infer):
if not infer:
returned_memories = []
for message in messages:
if message["role"] != "system":
message_embeddings = self.embedding_model.embed(message["content"], "add")
memory_id = self._create_memory(message["content"], message_embeddings, metadata)
returned_memories.append({"id": memory_id, "memory": message["content"], "event": "ADD"})
return returned_memories
parsed_messages = parse_messages(messages)
if self.custom_prompt:
@@ -305,15 +331,7 @@ class Memory(MemoryBase):
).model_dump(exclude={"score"})
# Add metadata if there are additional keys
excluded_keys = {
"user_id",
"agent_id",
"run_id",
"hash",
"data",
"created_at",
"updated_at",
}
excluded_keys = {"user_id", "agent_id", "run_id", "hash", "data", "created_at", "updated_at", "id"}
additional_metadata = {k: v for k, v in memory.payload.items() if k not in excluded_keys}
if additional_metadata:
memory_item["metadata"] = additional_metadata
@@ -376,6 +394,7 @@ class Memory(MemoryBase):
"data",
"created_at",
"updated_at",
"id",
}
all_memories = [
{
@@ -469,6 +488,7 @@ class Memory(MemoryBase):
"data",
"created_at",
"updated_at",
"id",
}
original_memories = [
@@ -655,4 +675,4 @@ class Memory(MemoryBase):
capture_event("mem0.reset", self)
def chat(self, query):
raise NotImplementedError("Chat function not implemented yet.")
raise NotImplementedError("Chat function not implemented yet.")
+1
View File
@@ -72,6 +72,7 @@ class VectorStoreFactory:
"vertex_ai_vector_search": "mem0.vector_stores.vertex_ai_vector_search.GoogleMatchingEngine",
"opensearch": "mem0.vector_stores.opensearch.OpenSearchDB",
"supabase": "mem0.vector_stores.supabase.Supabase",
"weaviate": "mem0.vector_stores.weaviate.Weaviate",
}
@classmethod
+144 -53
View File
@@ -1,5 +1,6 @@
import json
import logging
import re
from typing import List, Optional
from pydantic import BaseModel
@@ -12,6 +13,7 @@ try:
from azure.search.documents import SearchClient
from azure.search.documents.indexes import SearchIndexClient
from azure.search.documents.indexes.models import (
BinaryQuantizationCompression,
HnswAlgorithmConfiguration,
ScalarQuantizationCompression,
SearchField,
@@ -24,7 +26,7 @@ try:
from azure.search.documents.models import VectorizedQuery
except ImportError:
raise ImportError(
"The 'azure-search-documents' library is required. Please install it using 'pip install azure-search-documents==11.5.1'."
"The 'azure-search-documents' library is required. Please install it using 'pip install azure-search-documents==11.5.2'."
)
logger = logging.getLogger(__name__)
@@ -37,43 +39,82 @@ class OutputData(BaseModel):
class AzureAISearch(VectorStoreBase):
def __init__(self, service_name, collection_name, api_key, embedding_model_dims, use_compression):
"""Initialize the Azure Cognitive Search vector store.
def __init__(
self,
service_name,
collection_name,
api_key,
embedding_model_dims,
compression_type: Optional[str] = None,
use_float16: bool = False,
):
"""
Initialize the Azure AI Search vector store.
Args:
service_name (str): Azure Cognitive Search service name.
service_name (str): Azure AI Search service name.
collection_name (str): Index name.
api_key (str): API key for the Azure Cognitive Search service.
api_key (str): API key for the Azure AI Search service.
embedding_model_dims (int): Dimension of the embedding vector.
use_compression (bool): Use scalar quantization vector compression
compression_type (Optional[str]): Specifies the type of quantization to use.
Allowed values are None (no quantization), "scalar", or "binary".
use_float16 (bool): Whether to store vectors in half precision (Edm.Half) or full precision (Edm.Single).
(Note: This flag is preserved from the initial implementation per feedback.)
"""
self.index_name = collection_name
self.collection_name = collection_name
self.embedding_model_dims = embedding_model_dims
self.use_compression = use_compression
# If compression_type is None, treat it as "none".
self.compression_type = (compression_type or "none").lower()
self.use_float16 = use_float16
self.search_client = SearchClient(
endpoint=f"https://{service_name}.search.windows.net",
index_name=self.index_name,
credential=AzureKeyCredential(api_key),
)
self.index_client = SearchIndexClient(
endpoint=f"https://{service_name}.search.windows.net", credential=AzureKeyCredential(api_key)
endpoint=f"https://{service_name}.search.windows.net",
credential=AzureKeyCredential(api_key),
)
self.search_client._client._config.user_agent_policy.add_user_agent("mem0")
self.index_client._client._config.user_agent_policy.add_user_agent("mem0")
self.create_col() # create the collection / index
def create_col(self):
"""Create a new index in Azure Cognitive Search."""
vector_dimensions = self.embedding_model_dims # Set this to the number of dimensions in your vector
if self.use_compression:
"""Create a new index in Azure AI Search."""
# Determine vector type based on use_float16 setting.
if self.use_float16:
vector_type = "Collection(Edm.Half)"
compression_name = "myCompression"
compression_configurations = [ScalarQuantizationCompression(compression_name=compression_name)]
else:
vector_type = "Collection(Edm.Single)"
compression_name = None
compression_configurations = []
# Configure compression settings based on the specified compression_type.
compression_configurations = []
compression_name = None
if self.compression_type == "scalar":
compression_name = "myCompression"
# For SQ, rescoring defaults to True and oversampling defaults to 4.
compression_configurations = [
ScalarQuantizationCompression(
compression_name=compression_name
# rescoring defaults to True and oversampling defaults to 4
)
]
elif self.compression_type == "binary":
compression_name = "myCompression"
# For BQ, rescoring defaults to True and oversampling defaults to 10.
compression_configurations = [
BinaryQuantizationCompression(
compression_name=compression_name
# rescoring defaults to True and oversampling defaults to 10
)
]
# If no compression is desired, compression_configurations remains empty.
fields = [
SimpleField(name="id", type=SearchFieldDataType.String, key=True),
SimpleField(name="user_id", type=SearchFieldDataType.String, filterable=True),
@@ -82,8 +123,8 @@ class AzureAISearch(VectorStoreBase):
SearchField(
name="vector",
type=vector_type,
searchable=True,
vector_search_dimensions=vector_dimensions,
searchable=True,
vector_search_dimensions=self.embedding_model_dims,
vector_search_profile_name="my-vector-config",
),
SimpleField(name="payload", type=SearchFieldDataType.String, searchable=True),
@@ -91,7 +132,11 @@ class AzureAISearch(VectorStoreBase):
vector_search = VectorSearch(
profiles=[
VectorSearchProfile(name="my-vector-config", algorithm_configuration_name="my-algorithms-config")
VectorSearchProfile(
name="my-vector-config",
algorithm_configuration_name="my-algorithms-config",
compression_name=compression_name if self.compression_type != "none" else None
)
],
algorithms=[HnswAlgorithmConfiguration(name="my-algorithms-config")],
compressions=compression_configurations,
@@ -101,14 +146,16 @@ class AzureAISearch(VectorStoreBase):
def _generate_document(self, vector, payload, id):
document = {"id": id, "vector": vector, "payload": json.dumps(payload)}
# Extract additional fields if they exist
# Extract additional fields if they exist.
for field in ["user_id", "run_id", "agent_id"]:
if field in payload:
document[field] = payload[field]
return document
# Note: Explicit "insert" calls may later be decoupled from memory management decisions.
def insert(self, vectors, payloads=None, ids=None):
"""Insert vectors into the index.
"""
Insert vectors into the index.
Args:
vectors (List[List[float]]): List of vectors to insert.
@@ -116,61 +163,87 @@ class AzureAISearch(VectorStoreBase):
ids (List[str], optional): List of IDs corresponding to vectors.
"""
logger.info(f"Inserting {len(vectors)} vectors into index {self.index_name}")
documents = [
self._generate_document(vector, payload, id) for id, vector, payload in zip(ids, vectors, payloads)
self._generate_document(vector, payload, id)
for id, vector, payload in zip(ids, vectors, payloads)
]
self.search_client.upload_documents(documents)
response = self.search_client.upload_documents(documents)
for doc in response:
if not doc.get("status", False):
raise Exception(f"Insert failed for document {doc.get('id')}: {doc}")
return response
def _sanitize_key(self, key: str) -> str:
return re.sub(r"[^\w]", "", key)
def _build_filter_expression(self, filters):
filter_conditions = []
for key, value in filters.items():
# If the value is a string, add quotes
safe_key = self._sanitize_key(key)
if isinstance(value, str):
condition = f"{key} eq '{value}'"
safe_value = value.replace("'", "''")
condition = f"{safe_key} eq '{safe_value}'"
else:
condition = f"{key} eq {value}"
condition = f"{safe_key} eq {value}"
filter_conditions.append(condition)
# Use 'and' to join multiple conditions
filter_expression = " and ".join(filter_conditions)
return filter_expression
def search(self, query, limit=5, filters=None):
"""Search for similar vectors.
def search(self, query, limit=5, filters=None, vector_filter_mode="preFilter"):
"""
Search for similar vectors.
Args:
query (List[float]): Query vectors.
query (List[float]): Query vector.
limit (int, optional): Number of results to return. Defaults to 5.
filters (Dict, optional): Filters to apply to the search. Defaults to None.
vector_filter_mode (str): Determines whether filters are applied before or after the vector search.
Known values: "preFilter" (default) and "postFilter".
Returns:
list: Search results.
List[OutputData]: Search results.
"""
# Build filter expression
filter_expression = None
if filters:
filter_expression = self._build_filter_expression(filters)
vector_query = VectorizedQuery(vector=query, k_nearest_neighbors=limit, fields="vector")
search_results = self.search_client.search(vector_queries=[vector_query], filter=filter_expression, top=limit)
vector_query = VectorizedQuery(
vector=query, k_nearest_neighbors=limit, fields="vector"
)
search_results = self.search_client.search(
vector_queries=[vector_query],
filter=filter_expression,
top=limit,
vector_filter_mode=vector_filter_mode,
)
results = []
for result in search_results:
payload = json.loads(result["payload"])
results.append(OutputData(id=result["id"], score=result["@search.score"], payload=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.
"""
Delete a vector by ID.
Args:
vector_id (str): ID of the vector to delete.
"""
self.search_client.delete_documents(documents=[{"id": vector_id}])
response = self.search_client.delete_documents(documents=[{"id": vector_id}])
for doc in response:
if not doc.get("status", False):
raise Exception(f"Delete failed for document {vector_id}: {doc}")
logger.info(f"Deleted document with ID '{vector_id}' from index '{self.index_name}'.")
return response
def update(self, vector_id, vector=None, payload=None):
"""Update a vector and its payload.
"""
Update a vector and its payload.
Args:
vector_id (str): ID of the vector to update.
@@ -185,10 +258,15 @@ class AzureAISearch(VectorStoreBase):
document["payload"] = json_payload
for field in ["user_id", "run_id", "agent_id"]:
document[field] = payload.get(field)
self.search_client.merge_or_upload_documents(documents=[document])
response = self.search_client.merge_or_upload_documents(documents=[document])
for doc in response:
if not doc.get("status", False):
raise Exception(f"Update failed for document {vector_id}: {doc}")
return response
def get(self, vector_id) -> OutputData:
"""Retrieve a vector by ID.
"""
Retrieve a vector by ID.
Args:
vector_id (str): ID of the vector to retrieve.
@@ -200,35 +278,43 @@ class AzureAISearch(VectorStoreBase):
result = self.search_client.get_document(key=vector_id)
except ResourceNotFoundError:
return None
return OutputData(id=result["id"], score=None, payload=json.loads(result["payload"]))
return OutputData(
id=result["id"], score=None, payload=json.loads(result["payload"])
)
def list_cols(self) -> List[str]:
"""List all collections (indexes).
"""
List all collections (indexes).
Returns:
List[str]: List of index names.
"""
indexes = self.index_client.list_indexes()
return [index.name for index in indexes]
try:
names = self.index_client.list_index_names()
except AttributeError:
names = [index.name for index in self.index_client.list_indexes()]
return names
def delete_col(self):
"""Delete the index."""
self.index_client.delete_index(self.index_name)
def col_info(self):
"""Get information about the index.
"""
Get information about the index.
Returns:
Dict[str, Any]: Index information.
dict: Index information.
"""
index = self.index_client.get_index(self.index_name)
return {"name": index.name, "fields": index.fields}
def list(self, filters=None, limit=100):
"""List all vectors in the index.
"""
List all vectors in the index.
Args:
filters (Dict, optional): Filters to apply to the list.
filters (dict, optional): Filters to apply to the list.
limit (int, optional): Number of vectors to return. Defaults to 100.
Returns:
@@ -238,13 +324,18 @@ class AzureAISearch(VectorStoreBase):
if filters:
filter_expression = self._build_filter_expression(filters)
search_results = self.search_client.search(search_text="*", filter=filter_expression, top=limit)
search_results = self.search_client.search(
search_text="*", filter=filter_expression, top=limit
)
results = []
for result in search_results:
payload = json.loads(result["payload"])
results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload))
return [results]
results.append(
OutputData(
id=result["id"], score=result["@search.score"], payload=payload
)
)
return results
def __del__(self):
"""Close the search client when the object is deleted."""
+1
View File
@@ -21,6 +21,7 @@ class VectorStoreConfig(BaseModel):
"vertex_ai_vector_search": "GoogleMatchingEngineConfig",
"opensearch": "OpenSearchConfig",
"supabase": "SupabaseConfig",
"weaviate": "WeaviateConfig",
}
@model_validator(mode="after")
+1 -1
View File
@@ -1,6 +1,6 @@
import logging
import uuid
from typing import List, Optional, Dict, Any
from typing import List, Optional
from pydantic import BaseModel
+307
View File
@@ -0,0 +1,307 @@
import logging
import uuid
from typing import Dict, List, Mapping, Optional
from pydantic import BaseModel
try:
import weaviate
except ImportError:
raise ImportError(
"The 'weaviate' library is required. Please install it using 'pip install weaviate-client weaviate'."
)
import weaviate.classes.config as wvcc
from weaviate.classes.init import Auth
from weaviate.classes.query import Filter, MetadataQuery
from weaviate.util import get_valid_uuid
from mem0.vector_stores.base import VectorStoreBase
logger = logging.getLogger(__name__)
class OutputData(BaseModel):
id: str
score: float
payload: Dict
class Weaviate(VectorStoreBase):
def __init__(
self,
collection_name: str,
embedding_model_dims: int,
cluster_url: str = None,
auth_client_secret: str = None,
additional_headers: dict = None,
):
"""
Initialize the Weaviate vector store.
Args:
collection_name (str): Name of the collection/class in Weaviate.
embedding_model_dims (int): Dimensions of the embedding model.
client (WeaviateClient, optional): Existing Weaviate client instance. Defaults to None.
cluster_url (str, optional): URL for Weaviate server. Defaults to None.
auth_config (dict, optional): Authentication configuration for Weaviate. Defaults to None.
additional_headers (dict, optional): Additional headers for requests. Defaults to None.
"""
if "localhost" in cluster_url:
self.client = weaviate.connect_to_local(headers=additional_headers)
else:
self.client = weaviate.connect_to_wcs(
cluster_url=cluster_url,
auth_credentials=Auth.api_key(auth_client_secret),
headers=additional_headers,
)
self.collection_name = collection_name
self.create_col(embedding_model_dims)
def _parse_output(self, data: Dict) -> List[OutputData]:
"""
Parse the output data.
Args:
data (Dict): Output data.
Returns:
List[OutputData]: Parsed output data.
"""
keys = ["ids", "distances", "metadatas"]
values = []
for key in keys:
value = data.get(key, [])
if isinstance(value, list) and value and isinstance(value[0], list):
value = value[0]
values.append(value)
ids, distances, metadatas = values
max_length = max(len(v) for v in values if isinstance(v, list) and v is not None)
result = []
for i in range(max_length):
entry = OutputData(
id=ids[i] if isinstance(ids, list) and ids and i < len(ids) else None,
score=(distances[i] if isinstance(distances, list) and distances and i < len(distances) else None),
payload=(metadatas[i] if isinstance(metadatas, list) and metadatas and i < len(metadatas) else None),
)
result.append(entry)
return result
def create_col(self, vector_size, distance="cosine"):
"""
Create a new collection with the specified schema.
Args:
vector_size (int): Size of the vectors to be stored.
distance (str, optional): Distance metric for vector similarity. Defaults to "cosine".
"""
if self.client.collections.exists(self.collection_name):
logging.debug(f"Collection {self.collection_name} already exists. Skipping creation.")
return
properties = [
wvcc.Property(name="ids", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="hash", data_type=wvcc.DataType.TEXT),
wvcc.Property(
name="metadata",
data_type=wvcc.DataType.TEXT,
description="Additional metadata",
),
wvcc.Property(name="data", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="created_at", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="category", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="updated_at", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="user_id", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="agent_id", data_type=wvcc.DataType.TEXT),
wvcc.Property(name="run_id", data_type=wvcc.DataType.TEXT),
]
vectorizer_config = wvcc.Configure.Vectorizer.none()
vector_index_config = wvcc.Configure.VectorIndex.hnsw()
self.client.collections.create(
self.collection_name,
vectorizer_config=vectorizer_config,
vector_index_config=vector_index_config,
properties=properties,
)
def insert(self, vectors, payloads=None, ids=None):
"""
Insert vectors into a collection.
Args:
vectors (list): List of vectors to insert.
payloads (list, optional): List of payloads corresponding to vectors. Defaults to None.
ids (list, optional): List of IDs corresponding to vectors. Defaults to None.
"""
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
with self.client.batch.fixed_size(batch_size=100) as batch:
for idx, vector in enumerate(vectors):
object_id = ids[idx] if ids and idx < len(ids) else str(uuid.uuid4())
object_id = get_valid_uuid(object_id)
data_object = payloads[idx] if payloads and idx < len(payloads) else {}
# Ensure 'id' is not included in properties (it's used as the Weaviate object ID)
if "ids" in data_object:
del data_object["ids"]
batch.add_object(collection=self.collection_name, properties=data_object, uuid=object_id, vector=vector)
def search(self, query: List[float], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
"""
Search for similar vectors.
"""
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.hybrid(
query="",
vector=query,
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] # Include the id in the payload
results.append(
OutputData(
id=str(obj.uuid),
score=1
if obj.metadata.distance is None
else 1 - obj.metadata.distance, # Convert distance to score
payload=payload,
)
)
return results
def delete(self, vector_id):
"""
Delete a vector by ID.
Args:
vector_id: ID of the vector to delete.
"""
collection = self.client.collections.get(str(self.collection_name))
collection.data.delete_by_id(vector_id)
def update(self, vector_id, vector=None, payload=None):
"""
Update a vector and its payload.
Args:
vector_id: ID of the vector to update.
vector (list, optional): Updated vector. Defaults to None.
payload (dict, optional): Updated payload. Defaults to None.
"""
collection = self.client.collections.get(str(self.collection_name))
if payload:
collection.data.update(uuid=vector_id, properties=payload)
if vector:
existing_data = self.get(vector_id)
if existing_data:
existing_data = dict(existing_data)
if "id" in existing_data:
del existing_data["id"]
existing_payload: Mapping[str, str] = existing_data
collection.data.update(uuid=vector_id, properties=existing_payload, vector=vector)
def get(self, vector_id):
"""
Retrieve a vector by ID.
Args:
vector_id: ID of the vector to retrieve.
Returns:
dict: Retrieved vector and metadata.
"""
vector_id = get_valid_uuid(vector_id)
collection = self.client.collections.get(str(self.collection_name))
response = collection.query.fetch_object_by_id(
uuid=vector_id,
return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"],
)
# results = {}
# print("reponse",response)
# for obj in response.objects:
payload = response.properties.copy()
payload["id"] = str(response.uuid).split("'")[0]
results = OutputData(
id=str(response.uuid).split("'")[0],
score=1.0,
payload=payload,
)
return results
def list_cols(self):
"""
List all collections.
Returns:
list: List of collection names.
"""
collections = self.client.collections.list_all()
logger.debug(f"collections: {collections}")
print(f"collections: {collections}")
return {"collections": [{"name": col.name} for col in collections]}
def delete_col(self):
"""Delete a collection."""
self.client.collections.delete(self.collection_name)
def col_info(self):
"""
Get information about a collection.
Returns:
dict: Collection information.
"""
schema = self.client.collections.get(self.collection_name)
if schema:
return schema
return None
def list(self, filters=None, limit=100) -> List[OutputData]:
"""
List all vectors in a collection.
"""
collection = self.client.collections.get(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.fetch_objects(
limit=limit,
filters=combined_filter,
return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"],
)
results = []
for obj in response.objects:
payload = obj.properties.copy()
payload["id"] = str(obj.uuid).split("'")[0]
results.append(OutputData(id=str(obj.uuid).split("'")[0], score=1.0, payload=payload))
return [results]
Generated
+228 -327
View File
File diff suppressed because it is too large Load Diff
+4 -3
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "mem0ai"
version = "0.1.67"
version = "0.1.69"
description = "Long-term memory for AI Agents"
authors = ["Mem0 <founders@mem0.ai>"]
exclude = [
@@ -22,13 +22,14 @@ openai = "^1.33.0"
posthog = "^3.5.0"
pytz = "^2024.1"
sqlalchemy = "^2.0.31"
langchain-community = "^0.3.1"
langchain-neo4j = "^0.4.0"
neo4j = "^5.23.1"
rank-bm25 = "^0.2.2"
azure-search-documents = "^11.5.0"
psycopg2-binary = "^2.9.10"
[tool.poetry.extras]
graph = ["langchain-community", "neo4j", "rank-bm25"]
graph = ["langchain-neo4j", "neo4j", "rank-bm25"]
[tool.poetry.group.test.dependencies]
pytest = "^8.2.2"
+11 -53
View File
@@ -20,8 +20,10 @@ def mock_openai_client():
yield mock_client
def test_generate_response_without_tools(mock_openai_client):
config = BaseLlmConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P)
def test_generate_response(mock_openai_client):
config = BaseLlmConfig(
model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P
)
llm = AzureOpenAILLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -29,67 +31,21 @@ def test_generate_response_without_tools(mock_openai_client):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_openai_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
mock_openai_client.chat.completions.create.assert_called_once_with(
model=MODEL, messages=messages, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_openai_client):
config = BaseLlmConfig(model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P)
llm = AzureOpenAILLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_openai_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_openai_client.chat.completions.create.assert_called_once_with(
model=MODEL,
messages=messages,
temperature=TEMPERATURE,
max_tokens=MAX_TOKENS,
top_p=TOP_P,
tools=tools,
tool_choice="auto",
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
assert response == "I'm doing well, thank you for asking!"
@pytest.mark.parametrize(
@@ -128,4 +84,6 @@ def test_generate_with_http_proxies(default_headers):
api_version=None,
default_headers=default_headers,
)
mock_http_client.assert_called_once_with(proxies="http://testproxy.mem0.net:8000")
mock_http_client.assert_called_once_with(
proxies="http://testproxy.mem0.net:8000"
)
+32 -64
View File
@@ -16,33 +16,47 @@ def mock_deepseek_client():
def test_deepseek_llm_base_url():
# case1: default config with deepseek official base url
config = BaseLlmConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key")
config = BaseLlmConfig(
model="deepseek-chat",
temperature=0.7,
max_tokens=100,
top_p=1.0,
api_key="api_key",
)
llm = DeepSeekLLM(config)
assert str(llm.client.base_url) == "https://api.deepseek.com"
# case2: with env variable DEEPSEEK_API_BASE
provider_base_url = "https://api.provider.com/v1/"
os.environ["DEEPSEEK_API_BASE"] = provider_base_url
config = BaseLlmConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key")
config = BaseLlmConfig(
model="deepseek-chat",
temperature=0.7,
max_tokens=100,
top_p=1.0,
api_key="api_key",
)
llm = DeepSeekLLM(config)
assert str(llm.client.base_url) == provider_base_url
# case3: with config.deepseek_base_url
config_base_url = "https://api.config.com/v1/"
config = BaseLlmConfig(
model="deepseek-chat",
temperature=0.7,
max_tokens=100,
top_p=1.0,
api_key="api_key",
deepseek_base_url=config_base_url
model="deepseek-chat",
temperature=0.7,
max_tokens=100,
top_p=1.0,
api_key="api_key",
deepseek_base_url=config_base_url,
)
llm = DeepSeekLLM(config)
assert str(llm.client.base_url) == config_base_url
def test_generate_response_without_tools(mock_deepseek_client):
config = BaseLlmConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0)
def test_generate_response(mock_deepseek_client):
config = BaseLlmConfig(
model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0
)
llm = DeepSeekLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -50,64 +64,18 @@ def test_generate_response_without_tools(mock_deepseek_client):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_deepseek_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
mock_deepseek_client.chat.completions.create.assert_called_once_with(
model="deepseek-chat", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
model="deepseek-chat",
messages=messages,
temperature=0.7,
max_tokens=100,
top_p=1.0,
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_deepseek_client):
config = BaseLlmConfig(model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0)
llm = DeepSeekLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_deepseek_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_deepseek_client.chat.completions.create.assert_called_once_with(
model="deepseek-chat",
messages=messages,
temperature=0.7,
max_tokens=100,
top_p=1.0,
tools=tools,
tool_choice="auto"
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
+9 -79
View File
@@ -17,7 +17,9 @@ def mock_gemini_client():
def test_generate_response_without_tools(mock_gemini_client: Mock):
config = BaseLlmConfig(model="gemini-1.5-flash-latest", temperature=0.7, max_tokens=100, top_p=1.0)
config = BaseLlmConfig(
model="gemini-1.5-flash-latest", temperature=0.7, max_tokens=100, top_p=1.0
)
llm = GeminiLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -34,86 +36,14 @@ def test_generate_response_without_tools(mock_gemini_client: Mock):
mock_gemini_client.generate_content.assert_called_once_with(
contents=[
{"parts": "THIS IS A SYSTEM PROMPT. YOU MUST OBEY THIS: You are a helpful assistant.", "role": "user"},
{
"parts": "THIS IS A SYSTEM PROMPT. YOU MUST OBEY THIS: You are a helpful assistant.",
"role": "user",
},
{"parts": "Hello, how are you?", "role": "user"},
],
generation_config=GenerationConfig(temperature=0.7, max_output_tokens=100, top_p=1.0),
tools=None,
tool_config=content_types.to_tool_config(
{"function_calling_config": {"mode": "auto", "allowed_function_names": None}}
generation_config=GenerationConfig(
temperature=0.7, max_output_tokens=100, top_p=1.0
),
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_gemini_client: Mock):
config = BaseLlmConfig(model="gemini-1.5-flash-latest", temperature=0.7, max_tokens=100, top_p=1.0)
llm = GeminiLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_tool_call = Mock()
mock_tool_call.name = "add_memory"
mock_tool_call.args = {"data": "Today is a sunny day."}
mock_part = Mock()
mock_part.function_call = mock_tool_call
mock_part.text = "I've added the memory for you."
mock_content = Mock()
mock_content.parts = [mock_part]
mock_message = Mock()
mock_message.content = mock_content
mock_response = Mock(candidates=[mock_message])
mock_gemini_client.generate_content.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_gemini_client.generate_content.assert_called_once_with(
contents=[
{"parts": "THIS IS A SYSTEM PROMPT. YOU MUST OBEY THIS: You are a helpful assistant.", "role": "user"},
{"parts": "Add a new memory: Today is a sunny day.", "role": "user"},
],
generation_config=GenerationConfig(temperature=0.7, max_output_tokens=100, top_p=1.0),
tools=[
{
"function_declarations": [
{
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
}
]
}
],
tool_config=content_types.to_tool_config(
{"function_calling_config": {"mode": "auto", "allowed_function_names": None}}
),
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
+8 -52
View File
@@ -14,8 +14,10 @@ def mock_groq_client():
yield mock_client
def test_generate_response_without_tools(mock_groq_client):
config = BaseLlmConfig(model="llama3-70b-8192", temperature=0.7, max_tokens=100, top_p=1.0)
def test_generate_response(mock_groq_client):
config = BaseLlmConfig(
model="llama3-70b-8192", temperature=0.7, max_tokens=100, top_p=1.0
)
llm = GroqLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -23,64 +25,18 @@ def test_generate_response_without_tools(mock_groq_client):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_groq_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
mock_groq_client.chat.completions.create.assert_called_once_with(
model="llama3-70b-8192", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_groq_client):
config = BaseLlmConfig(model="llama3-70b-8192", temperature=0.7, max_tokens=100, top_p=1.0)
llm = GroqLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_groq_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_groq_client.chat.completions.create.assert_called_once_with(
model="llama3-70b-8192",
messages=messages,
temperature=0.7,
max_tokens=100,
top_p=1.0,
tools=tools,
tool_choice="auto",
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
assert response == "I'm doing well, thank you for asking!"
+11 -51
View File
@@ -13,17 +13,22 @@ def mock_litellm():
def test_generate_response_with_unsupported_model(mock_litellm):
config = BaseLlmConfig(model="unsupported-model", temperature=0.7, max_tokens=100, top_p=1)
config = BaseLlmConfig(
model="unsupported-model", temperature=0.7, max_tokens=100, top_p=1
)
llm = litellm.LiteLLM(config)
messages = [{"role": "user", "content": "Hello"}]
mock_litellm.supports_function_calling.return_value = False
with pytest.raises(ValueError, match="Model 'unsupported-model' in litellm does not support function calling."):
with pytest.raises(
ValueError,
match="Model 'unsupported-model' in LiteLLM does not support function calling.",
):
llm.generate_response(messages)
def test_generate_response_without_tools(mock_litellm):
def test_generate_response(mock_litellm):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
messages = [
@@ -32,7 +37,9 @@ def test_generate_response_without_tools(mock_litellm):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_litellm.completion.return_value = mock_response
mock_litellm.supports_function_calling.return_value = True
@@ -42,50 +49,3 @@ def test_generate_response_without_tools(mock_litellm):
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_litellm):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_litellm.completion.return_value = mock_response
mock_litellm.supports_function_calling.return_value = True
response = llm.generate_response(messages, tools=tools)
mock_litellm.completion.assert_called_once_with(
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1, tools=tools, tool_choice="auto"
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
+16 -51
View File
@@ -16,7 +16,9 @@ def mock_openai_client():
def test_openai_llm_base_url():
# case1: default config: with openai official base url
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key")
config = BaseLlmConfig(
model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
)
llm = OpenAILLM(config)
# Note: openai client will parse the raw base_url into a URL object, which will have a trailing slash
assert str(llm.client.base_url) == "https://api.openai.com/v1/"
@@ -24,7 +26,9 @@ def test_openai_llm_base_url():
# case2: with env variable OPENAI_API_BASE
provider_base_url = "https://api.provider.com/v1"
os.environ["OPENAI_API_BASE"] = provider_base_url
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key")
config = BaseLlmConfig(
model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key"
)
llm = OpenAILLM(config)
# Note: openai client will parse the raw base_url into a URL object, which will have a trailing slash
assert str(llm.client.base_url) == provider_base_url + "/"
@@ -32,14 +36,19 @@ def test_openai_llm_base_url():
# case3: with config.openai_base_url
config_base_url = "https://api.config.com/v1"
config = BaseLlmConfig(
model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0, api_key="api_key", openai_base_url=config_base_url
model="gpt-4o",
temperature=0.7,
max_tokens=100,
top_p=1.0,
api_key="api_key",
openai_base_url=config_base_url,
)
llm = OpenAILLM(config)
# Note: openai client will parse the raw base_url into a URL object, which will have a trailing slash
assert str(llm.client.base_url) == config_base_url + "/"
def test_generate_response_without_tools(mock_openai_client):
def test_generate_response(mock_openai_client):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0)
llm = OpenAILLM(config)
messages = [
@@ -48,7 +57,9 @@ def test_generate_response_without_tools(mock_openai_client):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_openai_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
@@ -57,49 +68,3 @@ def test_generate_response_without_tools(mock_openai_client):
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_openai_client):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0)
llm = OpenAILLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_openai_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_openai_client.chat.completions.create.assert_called_once_with(
model="gpt-4o", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0, tools=tools, tool_choice="auto"
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
+11 -52
View File
@@ -14,8 +14,13 @@ def mock_together_client():
yield mock_client
def test_generate_response_without_tools(mock_together_client):
config = BaseLlmConfig(model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0)
def test_generate_response(mock_together_client):
config = BaseLlmConfig(
model="mistralai/Mixtral-8x7B-Instruct-v0.1",
temperature=0.7,
max_tokens=100,
top_p=1.0,
)
llm = TogetherLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -23,64 +28,18 @@ def test_generate_response_without_tools(mock_together_client):
]
mock_response = Mock()
mock_response.choices = [Mock(message=Mock(content="I'm doing well, thank you for asking!"))]
mock_response.choices = [
Mock(message=Mock(content="I'm doing well, thank you for asking!"))
]
mock_together_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages)
mock_together_client.chat.completions.create.assert_called_once_with(
model="mistralai/Mixtral-8x7B-Instruct-v0.1", messages=messages, temperature=0.7, max_tokens=100, top_p=1.0
)
assert response == "I'm doing well, thank you for asking!"
def test_generate_response_with_tools(mock_together_client):
config = BaseLlmConfig(model="mistralai/Mixtral-8x7B-Instruct-v0.1", temperature=0.7, max_tokens=100, top_p=1.0)
llm = TogetherLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Add a new memory: Today is a sunny day."},
]
tools = [
{
"type": "function",
"function": {
"name": "add_memory",
"description": "Add a memory",
"parameters": {
"type": "object",
"properties": {"data": {"type": "string", "description": "Data to add to memory"}},
"required": ["data"],
},
},
}
]
mock_response = Mock()
mock_message = Mock()
mock_message.content = "I've added the memory for you."
mock_tool_call = Mock()
mock_tool_call.function.name = "add_memory"
mock_tool_call.function.arguments = '{"data": "Today is a sunny day."}'
mock_message.tool_calls = [mock_tool_call]
mock_response.choices = [Mock(message=mock_message)]
mock_together_client.chat.completions.create.return_value = mock_response
response = llm.generate_response(messages, tools=tools)
mock_together_client.chat.completions.create.assert_called_once_with(
model="mistralai/Mixtral-8x7B-Instruct-v0.1",
messages=messages,
temperature=0.7,
max_tokens=100,
top_p=1.0,
tools=tools,
tool_choice="auto",
)
assert response["content"] == "I've added the memory for you."
assert len(response["tool_calls"]) == 1
assert response["tool_calls"][0]["name"] == "add_memory"
assert response["tool_calls"][0]["arguments"] == {"data": "Today is a sunny day."}
assert response == "I'm doing well, thank you for asking!"
+1 -1
View File
@@ -51,7 +51,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"}
[{"role": "user", "content": "Test message"}], {"user_id": "test_user"}, {"user_id": "test_user"}, True
)
# Remove the conditional assertion for _add_to_graph
+542
View File
@@ -0,0 +1,542 @@
import json
from unittest.mock import Mock, patch, MagicMock, call
import pytest
from azure.core.exceptions import ResourceNotFoundError, HttpResponseError
# Import the AzureAISearch class and related models
from mem0.vector_stores.azure_ai_search import AzureAISearch, OutputData
from mem0.configs.vector_stores.azure_ai_search import AzureAISearchConfig
# Fixture to patch SearchClient and SearchIndexClient and create an instance of AzureAISearch.
@pytest.fixture
def mock_clients():
with patch("mem0.vector_stores.azure_ai_search.SearchClient") as MockSearchClient, \
patch("mem0.vector_stores.azure_ai_search.SearchIndexClient") as MockIndexClient, \
patch("mem0.vector_stores.azure_ai_search.AzureKeyCredential") as MockAzureKeyCredential:
# Create mocked instances for search and index clients.
mock_search_client = MockSearchClient.return_value
mock_index_client = MockIndexClient.return_value
# Mock the client._client._config.user_agent_policy.add_user_agent
mock_search_client._client = MagicMock()
mock_search_client._client._config.user_agent_policy.add_user_agent = Mock()
mock_index_client._client = MagicMock()
mock_index_client._client._config.user_agent_policy.add_user_agent = Mock()
# Stub required methods on search_client.
mock_search_client.upload_documents = Mock()
mock_search_client.upload_documents.return_value = [{"status": True, "id": "doc1"}]
mock_search_client.search = Mock()
mock_search_client.delete_documents = Mock()
mock_search_client.delete_documents.return_value = [{"status": True, "id": "doc1"}]
mock_search_client.merge_or_upload_documents = Mock()
mock_search_client.merge_or_upload_documents.return_value = [{"status": True, "id": "doc1"}]
mock_search_client.get_document = Mock()
mock_search_client.close = Mock()
# Stub required methods on index_client.
mock_index_client.create_or_update_index = Mock()
mock_index_client.list_indexes = Mock()
mock_index_client.list_index_names = Mock(return_value=["test-index"])
mock_index_client.delete_index = Mock()
# For col_info() we assume get_index returns an object with name and fields attributes.
fake_index = Mock()
fake_index.name = "test-index"
fake_index.fields = ["id", "vector", "payload", "user_id", "run_id", "agent_id"]
mock_index_client.get_index = Mock(return_value=fake_index)
mock_index_client.close = Mock()
yield mock_search_client, mock_index_client, MockAzureKeyCredential
@pytest.fixture
def azure_ai_search_instance(mock_clients):
mock_search_client, mock_index_client, _ = mock_clients
# Create an instance with dummy parameters.
instance = AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="test-api-key",
embedding_model_dims=3,
compression_type="binary", # testing binary quantization option
use_float16=True
)
# Return instance and clients for verification.
return instance, mock_search_client, mock_index_client
# --- Tests for AzureAISearchConfig ---
def test_config_validation_valid():
"""Test valid configurations are accepted."""
# Test minimal configuration
config = AzureAISearchConfig(
service_name="test-service",
api_key="test-api-key",
embedding_model_dims=768
)
assert config.collection_name == "mem0" # Default value
assert config.service_name == "test-service"
assert config.api_key == "test-api-key"
assert config.embedding_model_dims == 768
assert config.compression_type is None
assert config.use_float16 is False
# Test with all optional parameters
config = AzureAISearchConfig(
collection_name="custom-index",
service_name="test-service",
api_key="test-api-key",
embedding_model_dims=1536,
compression_type="scalar",
use_float16=True
)
assert config.collection_name == "custom-index"
assert config.compression_type == "scalar"
assert config.use_float16 is True
def test_config_validation_invalid_compression_type():
"""Test that invalid compression types are rejected."""
with pytest.raises(ValueError) as exc_info:
AzureAISearchConfig(
service_name="test-service",
api_key="test-api-key",
embedding_model_dims=768,
compression_type="invalid-type" # Not a valid option
)
assert "Invalid compression_type" in str(exc_info.value)
def test_config_validation_deprecated_use_compression():
"""Test that using the deprecated use_compression parameter raises an error."""
with pytest.raises(ValueError) as exc_info:
AzureAISearchConfig(
service_name="test-service",
api_key="test-api-key",
embedding_model_dims=768,
use_compression=True # Deprecated parameter
)
# Fix: Use a partial string match instead of exact match
assert "use_compression" in str(exc_info.value)
assert "no longer supported" in str(exc_info.value)
def test_config_validation_extra_fields():
"""Test that extra fields are rejected."""
with pytest.raises(ValueError) as exc_info:
AzureAISearchConfig(
service_name="test-service",
api_key="test-api-key",
embedding_model_dims=768,
unknown_parameter="value" # Extra field
)
assert "Extra fields not allowed" in str(exc_info.value)
assert "unknown_parameter" in str(exc_info.value)
# --- Tests for AzureAISearch initialization ---
def test_initialization(mock_clients):
"""Test AzureAISearch initialization with different parameters."""
mock_search_client, mock_index_client, mock_azure_key_credential = mock_clients
# Test with minimal parameters
instance = AzureAISearch(
service_name="test-service",
collection_name="test-index",
api_key="test-api-key",
embedding_model_dims=768
)
# Verify initialization parameters
assert instance.index_name == "test-index"
assert instance.collection_name == "test-index"
assert instance.embedding_model_dims == 768
assert instance.compression_type == "none" # Default when None is passed
assert instance.use_float16 is False
# Verify client creation
mock_azure_key_credential.assert_called_with("test-api-key")
assert "mem0" in mock_search_client._client._config.user_agent_policy.add_user_agent.call_args[0]
assert "mem0" in mock_index_client._client._config.user_agent_policy.add_user_agent.call_args[0]
# Verify index creation was called
mock_index_client.create_or_update_index.assert_called_once()
def test_initialization_with_compression_types(mock_clients):
"""Test initialization with different compression types."""
mock_search_client, mock_index_client, _ = mock_clients
# Test with scalar compression
instance = AzureAISearch(
service_name="test-service",
collection_name="scalar-index",
api_key="test-api-key",
embedding_model_dims=768,
compression_type="scalar"
)
assert instance.compression_type == "scalar"
# Capture the index creation call
args, _ = mock_index_client.create_or_update_index.call_args_list[-1]
index = args[0]
# Verify scalar compression was configured
assert hasattr(index.vector_search, 'compressions')
assert len(index.vector_search.compressions) > 0
assert "ScalarQuantizationCompression" in str(type(index.vector_search.compressions[0]))
# Test with binary compression
instance = AzureAISearch(
service_name="test-service",
collection_name="binary-index",
api_key="test-api-key",
embedding_model_dims=768,
compression_type="binary"
)
assert instance.compression_type == "binary"
# Capture the index creation call
args, _ = mock_index_client.create_or_update_index.call_args_list[-1]
index = args[0]
# Verify binary compression was configured
assert hasattr(index.vector_search, 'compressions')
assert len(index.vector_search.compressions) > 0
assert "BinaryQuantizationCompression" in str(type(index.vector_search.compressions[0]))
# Test with no compression
instance = AzureAISearch(
service_name="test-service",
collection_name="no-compression-index",
api_key="test-api-key",
embedding_model_dims=768,
compression_type=None
)
assert instance.compression_type == "none"
# Capture the index creation call
args, _ = mock_index_client.create_or_update_index.call_args_list[-1]
index = args[0]
# Verify no compression was configured
assert hasattr(index.vector_search, 'compressions')
assert len(index.vector_search.compressions) == 0
def test_initialization_with_float_precision(mock_clients):
"""Test initialization with different float precision settings."""
mock_search_client, mock_index_client, _ = mock_clients
# Test with half precision (float16)
instance = AzureAISearch(
service_name="test-service",
collection_name="float16-index",
api_key="test-api-key",
embedding_model_dims=768,
use_float16=True
)
assert instance.use_float16 is True
# Capture the index creation call
args, _ = mock_index_client.create_or_update_index.call_args_list[-1]
index = args[0]
# Find the vector field and check its type
vector_field = next((f for f in index.fields if f.name == "vector"), None)
assert vector_field is not None
assert "Edm.Half" in vector_field.type
# Test with full precision (float32)
instance = AzureAISearch(
service_name="test-service",
collection_name="float32-index",
api_key="test-api-key",
embedding_model_dims=768,
use_float16=False
)
assert instance.use_float16 is False
# Capture the index creation call
args, _ = mock_index_client.create_or_update_index.call_args_list[-1]
index = args[0]
# Find the vector field and check its type
vector_field = next((f for f in index.fields if f.name == "vector"), None)
assert vector_field is not None
assert "Edm.Single" in vector_field.type
# --- Tests for create_col method ---
def test_create_col(azure_ai_search_instance):
"""Test the create_col method creates an index with the correct configuration."""
instance, _, mock_index_client = azure_ai_search_instance
# create_col is called during initialization, so we check the call that was already made
mock_index_client.create_or_update_index.assert_called_once()
# Verify the index configuration
args, _ = mock_index_client.create_or_update_index.call_args
index = args[0]
# Check basic properties
assert index.name == "test-index"
assert len(index.fields) == 6 # id, user_id, run_id, agent_id, vector, payload
# Check that required fields are present
field_names = [f.name for f in index.fields]
assert "id" in field_names
assert "vector" in field_names
assert "payload" in field_names
assert "user_id" in field_names
assert "run_id" in field_names
assert "agent_id" in field_names
# Check that id is the key field
id_field = next(f for f in index.fields if f.name == "id")
assert id_field.key is True
# Check vector search configuration
assert index.vector_search is not None
assert len(index.vector_search.profiles) == 1
assert index.vector_search.profiles[0].name == "my-vector-config"
assert index.vector_search.profiles[0].algorithm_configuration_name == "my-algorithms-config"
# Check algorithms
assert len(index.vector_search.algorithms) == 1
assert index.vector_search.algorithms[0].name == "my-algorithms-config"
assert "HnswAlgorithmConfiguration" in str(type(index.vector_search.algorithms[0]))
# With binary compression and float16, we should have compression configuration
assert len(index.vector_search.compressions) == 1
assert index.vector_search.compressions[0].compression_name == "myCompression"
assert "BinaryQuantizationCompression" in str(type(index.vector_search.compressions[0]))
def test_create_col_scalar_compression(mock_clients):
"""Test creating a collection with scalar compression."""
mock_search_client, mock_index_client, _ = mock_clients
instance = AzureAISearch(
service_name="test-service",
collection_name="scalar-index",
api_key="test-api-key",
embedding_model_dims=768,
compression_type="scalar"
)
# Verify the index configuration
args, _ = mock_index_client.create_or_update_index.call_args
index = args[0]
# Check compression configuration
assert len(index.vector_search.compressions) == 1
assert index.vector_search.compressions[0].compression_name == "myCompression"
assert "ScalarQuantizationCompression" in str(type(index.vector_search.compressions[0]))
# Check profile references compression
assert index.vector_search.profiles[0].compression_name == "myCompression"
def test_create_col_no_compression(mock_clients):
"""Test creating a collection with no compression."""
mock_search_client, mock_index_client, _ = mock_clients
instance = AzureAISearch(
service_name="test-service",
collection_name="no-compression-index",
api_key="test-api-key",
embedding_model_dims=768,
compression_type=None
)
# Verify the index configuration
args, _ = mock_index_client.create_or_update_index.call_args
index = args[0]
# Check compression configuration - should be empty
assert len(index.vector_search.compressions) == 0
# Check profile doesn't reference compression
assert index.vector_search.profiles[0].compression_name is None
# --- Tests for insert method ---
def test_insert_single(azure_ai_search_instance):
"""Test inserting a single vector."""
instance, mock_search_client, _ = azure_ai_search_instance
vectors = [[0.1, 0.2, 0.3]]
payloads = [{"user_id": "user1", "run_id": "run1", "agent_id": "agent1"}]
ids = ["doc1"]
instance.insert(vectors, payloads, ids)
# Verify upload_documents was called correctly
mock_search_client.upload_documents.assert_called_once()
args, _ = mock_search_client.upload_documents.call_args
documents = args[0]
# Verify document structure
assert len(documents) == 1
assert documents[0]["id"] == "doc1"
assert documents[0]["vector"] == [0.1, 0.2, 0.3]
assert documents[0]["payload"] == json.dumps(payloads[0])
assert documents[0]["user_id"] == "user1"
assert documents[0]["run_id"] == "run1"
assert documents[0]["agent_id"] == "agent1"
def test_insert_multiple(azure_ai_search_instance):
"""Test inserting multiple vectors in one call."""
instance, mock_search_client, _ = azure_ai_search_instance
# Create multiple vectors
num_docs = 3
vectors = [[float(i)/10, float(i+1)/10, float(i+2)/10] for i in range(num_docs)]
payloads = [{"user_id": f"user{i}", "content": f"Test content {i}"} for i in range(num_docs)]
ids = [f"doc{i}" for i in range(num_docs)]
# Configure mock to return success for all documents
mock_search_client.upload_documents.return_value = [
{"status": True, "id": id_val} for id_val in ids
]
# Insert the documents
instance.insert(vectors, payloads, ids)
# Verify upload_documents was called with correct documents
mock_search_client.upload_documents.assert_called_once()
args, _ = mock_search_client.upload_documents.call_args
documents = args[0]
# Verify all documents were included
assert len(documents) == num_docs
# Check first document
assert documents[0]["id"] == "doc0"
assert documents[0]["vector"] == [0.0, 0.1, 0.2]
assert documents[0]["payload"] == json.dumps(payloads[0])
assert documents[0]["user_id"] == "user0"
# Check last document
assert documents[2]["id"] == "doc2"
assert documents[2]["vector"] == [0.2, 0.3, 0.4]
assert documents[2]["payload"] == json.dumps(payloads[2])
assert documents[2]["user_id"] == "user2"
def test_insert_with_error(azure_ai_search_instance):
"""Test insert when Azure returns an error for one or more documents."""
instance, mock_search_client, _ = azure_ai_search_instance
# Configure mock to return an error for one document
mock_search_client.upload_documents.return_value = [
{"status": False, "id": "doc1", "errorMessage": "Azure error"}
]
vectors = [[0.1, 0.2, 0.3]]
payloads = [{"user_id": "user1"}]
ids = ["doc1"]
# Insert should raise an exception
with pytest.raises(Exception) as exc_info:
instance.insert(vectors, payloads, ids)
assert "Insert failed for document doc1" in str(exc_info.value)
# Configure mock to return mixed success/failure for multiple documents
mock_search_client.upload_documents.return_value = [
{"status": True, "id": "doc1"},
{"status": False, "id": "doc2", "errorMessage": "Azure error"}
]
vectors = [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]
payloads = [{"user_id": "user1"}, {"user_id": "user2"}]
ids = ["doc1", "doc2"]
# Insert should raise an exception
with pytest.raises(Exception) as exc_info:
instance.insert(vectors, payloads, ids)
assert "Insert failed for document doc2" in str(exc_info.value)
def test_insert_with_missing_payload_fields(azure_ai_search_instance):
"""Test inserting with payloads missing some of the expected fields."""
instance, mock_search_client, _ = azure_ai_search_instance
vectors = [[0.1, 0.2, 0.3]]
payloads = [{"content": "Some content without user_id, run_id, or agent_id"}]
ids = ["doc1"]
instance.insert(vectors, payloads, ids)
# Verify upload_documents was called correctly
mock_search_client.upload_documents.assert_called_once()
args, _ = mock_search_client.upload_documents.call_args
documents = args[0]
# Verify document has payload but not the extra fields
assert len(documents) == 1
assert documents[0]["id"] == "doc1"
assert documents[0]["vector"] == [0.1, 0.2, 0.3]
assert documents[0]["payload"] == json.dumps(payloads[0])
assert "user_id" not in documents[0]
assert "run_id" not in documents[0]
assert "agent_id" not in documents[0]
def test_insert_with_http_error(azure_ai_search_instance):
"""Test insert when Azure client throws an HTTP error."""
instance, mock_search_client, _ = azure_ai_search_instance
# Configure mock to raise an HttpResponseError
mock_search_client.upload_documents.side_effect = HttpResponseError("Azure service error")
vectors = [[0.1, 0.2, 0.3]]
payloads = [{"user_id": "user1"}]
ids = ["doc1"]
# Insert should propagate the HTTP error
with pytest.raises(HttpResponseError) as exc_info:
instance.insert(vectors, payloads, ids)
assert "Azure service error" in str(exc_info.value)
# --- Tests for search method ---
def test_search_basic(azure_ai_search_instance):
"""Test basic vector search without filters."""
instance, mock_search_client, _ = azure_ai_search_instance
# Configure mock to return search results
mock_search_client.search.return_value = [
{
"id": "doc1",
"@search.score": 0.95,
"payload": json.dumps({"content": "Test content"})
}
]
# Search with a vector
query_vector = [0.1, 0.2, 0.3]
results = instance.search(query_vector, limit=5)
# Verify search was called correctly
mock_search_client.search.assert_called_once()
_, kwargs = mock_search_client.search.call_args
# Check parameters
assert len(kwargs["vector_queries"]) == 1
assert kwargs["vector_queries"][0].vector == query_vector
assert kwargs["vector_queries"][0].k_nearest_neighbors == 5
assert kwargs["vector_queries"][0].fields == "vector"
assert kwargs["filter"] is None # No filters
assert kwargs["top"] == 5
assert kwargs["vector_filter_mode"] == "preFilter" # Default mode
# Check results
assert len(results) == 1
assert results[0].id == "doc1"
assert results[0].score == 0.95
assert results[0].payload == {"content": "Test content"}
+220
View File
@@ -0,0 +1,220 @@
import os
import uuid
import httpx
import unittest
from unittest.mock import MagicMock, patch
import dotenv
import weaviate
from weaviate.classes.query import MetadataQuery, Filter
from weaviate.exceptions import UnexpectedStatusCodeException
from mem0.vector_stores.weaviate import Weaviate, OutputData
class TestWeaviateDB(unittest.TestCase):
@classmethod
def setUpClass(cls):
dotenv.load_dotenv()
cls.original_env = {
'WEAVIATE_CLUSTER_URL': os.getenv('WEAVIATE_CLUSTER_URL', 'http://localhost:8080'),
'WEAVIATE_API_KEY': os.getenv('WEAVIATE_API_KEY', 'test_api_key'),
}
os.environ['WEAVIATE_CLUSTER_URL'] = 'http://localhost:8080'
os.environ['WEAVIATE_API_KEY'] = 'test_api_key'
def setUp(self):
self.client_mock = MagicMock(spec=weaviate.WeaviateClient)
self.client_mock.collections = MagicMock()
self.client_mock.collections.exists.return_value = False
self.client_mock.collections.create.return_value = None
self.client_mock.collections.delete.return_value = None
patcher = patch('mem0.vector_stores.weaviate.weaviate.connect_to_local', return_value=self.client_mock)
self.mock_weaviate = patcher.start()
self.addCleanup(patcher.stop)
self.weaviate_db = Weaviate(
collection_name="test_collection",
embedding_model_dims=1536,
cluster_url=os.getenv('WEAVIATE_CLUSTER_URL'),
auth_client_secret=os.getenv('WEAVIATE_API_KEY'),
additional_headers={"X-OpenAI-Api-Key": "test_key"},
)
self.client_mock.reset_mock()
@classmethod
def tearDownClass(cls):
for key, value in cls.original_env.items():
if value is not None:
os.environ[key] = value
else:
os.environ.pop(key, None)
def tearDown(self):
self.client_mock.reset_mock()
def test_create_col(self):
self.client_mock.collections.exists.return_value = False
self.weaviate_db.create_col(vector_size=1536)
self.client_mock.collections.create.assert_called_once()
self.client_mock.reset_mock()
self.client_mock.collections.exists.return_value = True
self.weaviate_db.create_col(vector_size=1536)
self.client_mock.collections.create.assert_not_called()
def test_insert(self):
self.client_mock.batch = MagicMock()
self.client_mock.batch.fixed_size.return_value.__enter__.return_value = MagicMock()
self.client_mock.collections.get.return_value.data.insert_many.return_value = {
"results": [{"id": "id1"}, {"id": "id2"}]
}
vectors = [[0.1] * 1536, [0.2] * 1536]
payloads = [{"key1": "value1"}, {"key2": "value2"}]
ids = [str(uuid.uuid4()), str(uuid.uuid4())]
results = self.weaviate_db.insert(vectors=vectors, payloads=payloads, ids=ids)
def test_get(self):
valid_uuid = str(uuid.uuid4())
mock_response = MagicMock()
mock_response.properties = {
"hash": "abc123",
"created_at": "2025-03-08T12:00:00Z",
"updated_at": "2025-03-08T13:00:00Z",
"user_id": "user_123",
"agent_id": "agent_456",
"run_id": "run_789",
"data": {"key": "value"},
"category": "test",
}
mock_response.uuid = valid_uuid
self.client_mock.collections.get.return_value.query.fetch_object_by_id.return_value = mock_response
result = self.weaviate_db.get(vector_id=valid_uuid)
assert result.id == valid_uuid
expected_payload = mock_response.properties.copy()
expected_payload["id"] = valid_uuid
assert result.payload == expected_payload
def test_get_not_found(self):
mock_response = httpx.Response(status_code=404, json={"error": "Not found"})
self.client_mock.collections.get.return_value.data.get_by_id.side_effect = UnexpectedStatusCodeException(
"Not found", mock_response
)
def test_search(self):
mock_objects = [
{
"uuid": "id1",
"properties": {"key1": "value1"},
"metadata": {"distance": 0.2}
}
]
mock_response = MagicMock()
mock_response.objects = []
for obj in mock_objects:
mock_obj = MagicMock()
mock_obj.uuid = obj["uuid"]
mock_obj.properties = obj["properties"]
mock_obj.metadata = MagicMock()
mock_obj.metadata.distance = obj["metadata"]["distance"]
mock_response.objects.append(mock_obj)
mock_hybrid = MagicMock()
self.client_mock.collections.get.return_value.query.hybrid = mock_hybrid
mock_hybrid.return_value = mock_response
query_vector = [0.1] * 1536
results = self.weaviate_db.search(query=query_vector, limit=5)
mock_hybrid.assert_called_once()
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, "id1")
self.assertEqual(results[0].score, 0.8)
def test_delete(self):
self.weaviate_db.delete(vector_id="id1")
self.client_mock.collections.get.return_value.data.delete_by_id.assert_called_once_with("id1")
def test_list(self):
mock_objects = []
mock_obj1 = MagicMock()
mock_obj1.uuid = "id1"
mock_obj1.properties = {"key1": "value1"}
mock_objects.append(mock_obj1)
mock_obj2 = MagicMock()
mock_obj2.uuid = "id2"
mock_obj2.properties = {"key2": "value2"}
mock_objects.append(mock_obj2)
mock_response = MagicMock()
mock_response.objects = mock_objects
mock_fetch = MagicMock()
self.client_mock.collections.get.return_value.query.fetch_objects = mock_fetch
mock_fetch.return_value = mock_response
results = self.weaviate_db.list(limit=10)
mock_fetch.assert_called_once()
# Verify results
self.assertEqual(len(results), 1)
self.assertEqual(len(results[0]), 2)
self.assertEqual(results[0][0].id, "id1")
self.assertEqual(results[0][0].payload["key1"], "value1")
self.assertEqual(results[0][1].id, "id2")
self.assertEqual(results[0][1].payload["key2"], "value2")
def test_list_cols(self):
mock_collection1 = MagicMock()
mock_collection1.name = "collection1"
mock_collection2 = MagicMock()
mock_collection2.name = "collection2"
self.client_mock.collections.list_all.return_value = [mock_collection1, mock_collection2]
result = self.weaviate_db.list_cols()
expected = {"collections": [{"name": "collection1"}, {"name": "collection2"}]}
assert result == expected
self.client_mock.collections.list_all.assert_called_once()
def test_delete_col(self):
self.weaviate_db.delete_col()
self.client_mock.collections.delete.assert_called_once_with("test_collection")
if __name__ == '__main__':
unittest.main()