Compare commits
13 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| d7a26bd0c3 | |||
| e25dc4b504 | |||
| dab3349990 | |||
| 6db87e8d07 | |||
| faf811ee2d | |||
| ee80a43810 | |||
| 4be426f762 | |||
| ba9c61938b | |||
| 65f826e064 | |||
| b43363cdf3 | |||
| 89e786a88e | |||
| 2d5062bd40 | |||
| b89628322d |
@@ -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` |
|
||||
@@ -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
@@ -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"
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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>
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
};
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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/`,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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}")
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -21,6 +21,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"vertex_ai_vector_search": "GoogleMatchingEngineConfig",
|
||||
"opensearch": "OpenSearchConfig",
|
||||
"supabase": "SupabaseConfig",
|
||||
"weaviate": "WeaviateConfig",
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import logging
|
||||
import uuid
|
||||
from typing import List, Optional, Dict, Any
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+4
-3
@@ -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"
|
||||
|
||||
@@ -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
@@ -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."}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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"}
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user