Compare commits

...

26 Commits

Author SHA1 Message Date
Dev-Khant ff30cb8ddd version bump -> 0.1.74 2025-03-21 13:06:40 +05:30
Parshva Daftari 2e853c3d22 Updated VDB Docs (#2409) 2025-03-20 23:47:57 +05:30
Dev Khant 3cc7013fde fix pinecone (#2414) 2025-03-20 23:47:09 +05:30
Dev Khant 8e6a08aa83 Support for hybrid search in Azure AI vector store (#2408)
Co-authored-by: Deshraj Yadav <deshrajdry@gmail.com>
2025-03-20 22:57:00 +05:30
Wonbin Kim 8b9a8e5825 URGENT Hotfix - update default Elasticsearch search query (#2413) 2025-03-20 20:50:18 +05:30
Dev Khant afc630272d bump version -> 0.1.73 (#2412) 2025-03-20 19:30:29 +05:30
Parshva Daftari e33008e3a4 Add: Pinecone integration (#2395) 2025-03-20 12:57:32 +05:30
Mauricio A 7b516328a8 Feature/fix opensearch vector mapping (#2399) 2025-03-20 09:37:57 +05:30
Dev Khant 6d5889d98f version bump -> 0.1.72 (#2405) 2025-03-20 00:10:27 +05:30
Parshva Daftari ee66e0c954 Reverting the tools commit (#2404) 2025-03-20 00:09:00 +05:30
Prateek Chhikara 1aed611539 Added graph memory (#2403) 2025-03-19 09:51:15 -07:00
Saket Aryan 6c2b131d6e Added Feedback in SDK (#2393) 2025-03-19 09:11:45 -07:00
Gaurav Agerwala 2ffe9922f3 Added support for Ollama in TS SDK (#2345)
Co-authored-by: Dev Khant <devkhant24@gmail.com>
2025-03-19 21:38:19 +05:30
Dev Khant 540ada489b version bump -> 0.1.71 (#2402) 2025-03-19 21:36:06 +05:30
Dev Khant 65cffa0369 Fix: made tools support for graph (#2400) 2025-03-19 21:29:36 +05:30
Dev Khant 9f937943ba Doc: update oss quickstart page (#2401) 2025-03-19 17:27:09 +05:30
Dev-Khant 51a68bf7c5 Doc: update azure ai vector store 2025-03-18 14:32:12 +05:30
Dev-Khant 92541d8955 Doc: azure ai vector search 2025-03-18 14:17:56 +05:30
Dev Khant 0e0be18ecc Fix azure ai vector store (#2396) 2025-03-18 14:13:19 +05:30
Wonbin Kim 66d3f9b93c Support Custom Prompt for Memory Action Decision (#2371) 2025-03-18 10:43:01 +05:30
Wonbin Kim b8f40f728f Support Custom Search Query for Elasticsearch (#2372) 2025-03-18 10:34:34 +05:30
Prateek Chhikara 00a2ea9ff0 Added export instructions to docs (#2394) 2025-03-17 17:41:56 -07:00
Prateek Chhikara 9545836469 Added docs for add-v2 (#2381) 2025-03-17 15:39:17 -07:00
Saket Aryan 3acd9e20da Fix Redis Search (#2392) 2025-03-17 15:30:40 -07:00
Dev Khant d48ecd52ef update poetry lock file (#2391) 2025-03-18 01:11:05 +05:30
Saket Aryan 2fbea7705b Add Intercom to Docs (#2390) 2025-03-17 12:37:56 -07:00
86 changed files with 3843 additions and 1023 deletions
+1 -1
View File
@@ -13,7 +13,7 @@ install:
install_all:
poetry install
poetry run pip install groq together boto3 litellm ollama chromadb weaviate weaviate-client sentence_transformers vertexai \
google-generativeai elasticsearch opensearch-py vecs
google-generativeai elasticsearch opensearch-py vecs pinecone pinecone-text
# Format code with ruff
format:
@@ -1,5 +1,3 @@
# Azure AI Search
[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
@@ -17,8 +15,7 @@ config = {
"service_name": "ai-search-test",
"api_key": "*****",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"compression_type": "none"
"embedding_model_dims": 1536
}
}
}
@@ -33,18 +30,9 @@ messages = [
m.add(messages, user_id="alice", metadata={"category": "movies"})
```
## Advanced Usage
## Using binary compression for large vector collections
```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",
@@ -60,6 +48,24 @@ config = {
}
```
## Using hybrid search
```python
config = {
"vector_store": {
"provider": "azure_ai_search",
"config": {
"service_name": "ai-search-test",
"api_key": "*****",
"collection_name": "mem0",
"embedding_model_dims": 1536,
"hybrid_search": True,
"vector_filter_mode": "postFilter"
}
}
}
```
## Configuration Parameters
| Parameter | Description | Default Value | Options |
@@ -70,6 +76,8 @@ config = {
| `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` |
| `vector_filter_mode` | Vector filter mode to use | `preFilter` | `postFilter`, `preFilter` |
| `hybrid_search` | Use hybrid search | `False` | `True`, `False` |
## Notes on Configuration Options
@@ -54,6 +54,7 @@ Let's see the available parameters for the `elasticsearch` config:
| `password` | Password for basic authentication | `None` |
| `verify_certs` | Whether to verify SSL certificates | `True` |
| `auto_create_index` | Whether to automatically create the index | `True` |
| `custom_search_query` | Function returning a custom search query | `None` |
### Features
@@ -62,3 +63,46 @@ Let's see the available parameters for the `elasticsearch` config:
- Multiple authentication methods (Basic Auth, API Key)
- Automatic index creation with optimized mappings for vector search
- Memory isolation through payload filtering
- Custom search query function to customize the search query
### Custom Search Query
The `custom_search_query` parameter allows you to customize the search query when `Memory.search` is called.
__Example__
```python
import os
from typing import List, Optional, Dict
from mem0 import Memory
def custom_search_query(query: List[float], limit: int, filters: Optional[Dict]) -> Dict:
return {
"knn": {
"field": "vector",
"query_vector": query,
"k": limit,
"num_candidates": limit * 2
}
}
os.environ["OPENAI_API_KEY"] = "sk-xx"
config = {
"vector_store": {
"provider": "elasticsearch",
"config": {
"collection_name": "mem0",
"host": "localhost",
"port": 9200,
"embedding_model_dims": 1536,
"custom_search_query": custom_search_query
}
}
}
```
It should be a function that takes the following parameters:
- `query`: a query vector used in `Memory.search`
- `limit`: a number of results used in `Memory.search`
- `filters`: a dictionary of key-value pairs used in `Memory.search`. You can add custom pairs for the custom search query.
The function should return a query body for the Elasticsearch search API.
@@ -0,0 +1,92 @@
[Pinecone](https://www.pinecone.io/) is a fully managed vector database designed for machine learning applications, offering high performance vector search with low latency at scale. It's particularly well-suited for semantic search, recommendation systems, and other AI-powered applications.
> **Note**: Before configuring Pinecone, you need to select an embedding model (e.g., OpenAI, Cohere, or custom models) and ensure the `embedding_model_dims` in your config matches your chosen model's dimensions. For example, OpenAI's text-embedding-ada-002 uses 1536 dimensions.
### Usage
```python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "sk-xx"
os.environ["PINECONE_API_KEY"] = "your-api-key"
# Example using serverless configuration
config = {
"vector_store": {
"provider": "pinecone",
"config": {
"collection_name": "testing",
"embedding_model_dims": 1536, # Matches OpenAI's text-embedding-3-small
"serverless_config": {
"cloud": "aws", # Choose between 'aws' or 'gcp' or 'azure'
"region": "us-east-1"
},
"metric": "cosine"
}
}
}
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": "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
Here are the parameters available for configuring Pinecone:
| Parameter | Description | Default Value |
| --- | --- | --- |
| `collection_name` | Name of the index/collection | Required |
| `embedding_model_dims` | Dimensions of the embedding model (must match your chosen embedding model) | Required |
| `client` | Existing Pinecone client instance | `None` |
| `api_key` | API key for Pinecone | Environment variable: `PINECONE_API_KEY` |
| `environment` | Pinecone environment | `None` |
| `serverless_config` | Configuration for serverless deployment (AWS or GCP or Azure) | `None` |
| `pod_config` | Configuration for pod-based deployment | `None` |
| `hybrid_search` | Whether to enable hybrid search | `False` |
| `metric` | Distance metric for vector similarity | `"cosine"` |
| `batch_size` | Batch size for operations | `100` |
> **Important**: You must choose either `serverless_config` or `pod_config` for your deployment, but not both.
#### Serverless Config Example
```python
config = {
"vector_store": {
"provider": "pinecone",
"config": {
"collection_name": "memory_index",
"embedding_model_dims": 1536, # For OpenAI's text-embedding-3-small
"serverless_config": {
"cloud": "aws", # or "gcp" or "azure"
"region": "us-east-1" # Choose appropriate region
}
}
}
}
```
#### Pod Config Example
```python
config = {
"vector_store": {
"provider": "pinecone",
"config": {
"collection_name": "memory_index",
"embedding_model_dims": 1536, # For OpenAI's text-embedding-ada-002
"pod_config": {
"environment": "gcp-starter",
"replicas": 1,
"pod_type": "starter"
}
}
}
}
```
+1
View File
@@ -19,6 +19,7 @@ See the list of supported vector databases below.
<Card title="Chroma" href="/components/vectordbs/dbs/chroma"></Card>
<Card title="Pgvector" href="/components/vectordbs/dbs/pgvector"></Card>
<Card title="Milvus" href="/components/vectordbs/dbs/milvus"></Card>
<Card title="Pinecone" href="/components/vectordbs/dbs/pinecone"></Card>
<Card title="Azure AI Search" href="/components/vectordbs/dbs/azure_ai_search"></Card>
<Card title="Redis" href="/components/vectordbs/dbs/redis"></Card>
<Card title="Elasticsearch" href="/components/vectordbs/dbs/elasticsearch"></Card>
+10 -2
View File
@@ -47,6 +47,7 @@
"pages": [
"features/platform-overview",
"features/advanced-retrieval",
"features/contextual-add",
"features/multimodal-support",
"features/selective-memory",
"features/custom-categories",
@@ -54,7 +55,9 @@
"features/direct-import",
"features/async-client",
"features/memory-export",
"features/webhooks"
"features/webhooks",
"features/graph-memory",
"features/feedback-mechanism"
]
}
]
@@ -71,7 +74,8 @@
"icon": "wrench",
"pages": [
"features/openai_compatibility",
"features/custom-prompts",
"features/custom-fact-extraction-prompt",
"features/custom-update-memory-prompt",
"open-source/multimodal-support",
"open-source/features/rest-api"
]
@@ -125,6 +129,7 @@
"components/vectordbs/dbs/chroma",
"components/vectordbs/dbs/pgvector",
"components/vectordbs/dbs/milvus",
"components/vectordbs/dbs/pinecone",
"components/vectordbs/dbs/azure_ai_search",
"components/vectordbs/dbs/redis",
"components/vectordbs/dbs/elasticsearch",
@@ -336,6 +341,9 @@
"posthog": {
"apiKey": "phc_hgJkUVJFYtmaJqrvf6CYN67TIQ8yhXAkWzUn9AMU4yX",
"apiHost": "https://mango.mem0.ai"
},
"intercom": {
"appId": "jjv2r0tt"
}
}
}
+205
View File
@@ -0,0 +1,205 @@
---
title: Contextual Add (ADD v2)
icon: "square-plus"
iconType: "solid"
---
Mem0 now supports an contextual add version (v2). To use it, set `version="v2"` during the add call. The default version is v1, which is deprecated now. We recommend migrating to `v2` for new applications.
## Key Differences Between v1 and v2
### Version 1 (Legacy)
In v1 (default), users needed to pass either the entire conversation history or past k messages with each new message to generate properly contextualized memories. This approach required:
- Manually tracking and sending previous messages using a sliding window approach
- Increased payload sizes as conversations grew longer, requiring careful window size management
<CodeGroup>
```python Python
# First interaction
messages1 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."}
]
client.add(messages1, user_id="alex")
# Second interaction - must include previous messages for context
messages2 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."},
{"role": "user", "content": "I like to eat sushi, and yesterday I went to Sunnyvale to eat sushi with my friends."},
{"role": "assistant", "content": "Sushi is really a tasty choice. What did you do this weekend?"}
]
client.add(messages2, user_id="alex")
```
```javascript JavaScript
// First interaction
const messages1 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."}
];
client.add(messages1, { user_id: "alex" })
.then(response => console.log(response))
.catch(error => console.error(error));
// Second interaction - must include previous messages for context
const messages2 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."},
{"role": "user", "content": "I like to eat sushi, and yesterday I went to Sunnyvale to eat sushi with my friends."},
{"role": "assistant", "content": "Sushi is really a tasty choice. What did you do this weekend?"}
];
client.add(messages2, { user_id: "alex" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
</CodeGroup>
### Version 2 (Recommended)
In v2, Mem0 automatically manages conversation context. Users only need to send new messages, and the system will:
- Automatically retrieve relevant conversation history
- Generate properly contextualized memories
- Reduce payload sizes and simplify integration
<CodeGroup>
```python Python
# First interaction
messages1 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."}
]
client.add(messages1, user_id="alex", version="v2")
# Second interaction - only need to send new messages
messages2 = [
{"role": "user", "content": "I like to eat sushi, and yesterday I went to Sunnyvale to eat sushi with my friends."},
{"role": "assistant", "content": "Sushi is really a tasty choice. What did you do this weekend?"}
]
client.add(messages2, user_id="alex", version="v2")
```
```javascript JavaScript
// First interaction
const messages1 = [
{"role": "user", "content": "Hi, I'm Alex and I live in San Francisco."},
{"role": "assistant", "content": "Hello Alex! Nice to meet you. San Francisco is a beautiful city."}
];
client.add(messages1, { user_id: "alex", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
// Second interaction - only need to send new messages
const messages2 = [
{"role": "user", "content": "I like to eat sushi, and yesterday I went to Sunnyvale to eat sushi with my friends."},
{"role": "assistant", "content": "Sushi is really a tasty choice. What did you do this weekend?"}
];
client.add(messages2, { user_id: "alex", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
</CodeGroup>
## Benefits of Using v2
1. **Simplified Integration**: No need to track and manage conversation history
2. **Reduced Payload Size**: Only send new messages, not the entire conversation
3. **Improved Memory Quality**: Automatic context retrieval ensures better memory generation
## Understanding ID Parameters in v2
When using contextual add v2, you have different options for how to organize and retrieve memories:
### Using Only `user_id`
When you provide only a `user_id`:
- Memories are associated with this user's long-term memory store
- The system will automatically retrieve relevant context from all of the user's previous conversations
- These memories persist indefinitely across all of the user's sessions
- Ideal for maintaining persistent user information (preferences, personal details, etc.)
<CodeGroup>
```python Python
# Adding to long-term user memory
messages = [
{"role": "user", "content": "I'm allergic to peanuts and shellfish."},
{"role": "assistant", "content": "I've noted your allergies to peanuts and shellfish."}
]
client.add(messages, user_id="alex", version="v2")
```
```javascript JavaScript
// Adding to long-term user memory
const messages = [
{"role": "user", "content": "I'm allergic to peanuts and shellfish."},
{"role": "assistant", "content": "I've noted your allergies to peanuts and shellfish."}
];
client.add(messages, { user_id: "alex", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
</CodeGroup>
### Using `user_id` with `run_id`
When you provide both `user_id` and `run_id`:
- Memories are associated with a specific conversation session or interaction
- The system will retrieve context primarily from this specific session
- These memories are still tied to the user but are organized by the specific session
- Ideal for maintaining context within a specific conversation flow or task
- Helps prevent context from different conversations from interfering with each other
<CodeGroup>
```python Python
# Adding to a specific conversation session
messages = [
{"role": "user", "content": "For this trip to Paris, I want to focus on art museums."},
{"role": "assistant", "content": "Great! I'll help you plan your Paris trip with a focus on art museums."}
]
client.add(messages, user_id="alex", run_id="paris-trip-2024", version="v2")
# Later in the same conversation session
messages2 = [
{"role": "user", "content": "I'd like to visit the Louvre on Monday."},
{"role": "assistant", "content": "The Louvre is a great choice for Monday. Would you like information about opening hours?"}
]
client.add(messages2, user_id="alex", run_id="paris-trip-2024", version="v2")
```
```javascript JavaScript
// Adding to a specific conversation session
const messages = [
{"role": "user", "content": "For this trip to Paris, I want to focus on art museums."},
{"role": "assistant", "content": "Great! I'll help you plan your Paris trip with a focus on art museums."}
];
client.add(messages, { user_id: "alex", run_id: "paris-trip-2024", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
// Later in the same conversation session
const messages2 = [
{"role": "user", "content": "I'd like to visit the Louvre on Monday."},
{"role": "assistant", "content": "The Louvre is a great choice for Monday. Would you like information about opening hours?"}
];
client.add(messages2, { user_id: "alex", run_id: "paris-trip-2024", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
</CodeGroup>
Using `run_id` helps you organize memories into logical sessions or tasks, making it easier to maintain context for specific interactions while still associating everything with the user's overall profile.
If you have any questions, please feel free to reach out to us using one of the following methods:
<Snippet file="get-help.mdx" />
@@ -1,25 +1,25 @@
---
title: Custom Prompts
description: 'Enhance your product experience by adding custom prompts tailored to your needs'
title: Custom Fact Extraction Prompt
description: 'Enhance your product experience by adding custom fact extraction prompt tailored to your needs'
icon: "pencil"
iconType: "solid"
---
## Introduction to Custom Prompts
## Introduction to Custom Fact Extraction Prompt
Custom prompts allow you to tailor the behavior of your Mem0 instance to specific use cases or domains.
By defining a custom prompt, you can control how information is extracted, processed, and stored in your memory system.
Custom fact extraction prompt allow you to tailor the behavior of your Mem0 instance to specific use cases or domains.
By defining it, you can control how information is extracted from the user's message.
To create an effective custom prompt:
To create an effective custom fact extraction prompt:
1. Be specific about the information to extract.
2. Provide few-shot examples to guide the LLM.
3. Ensure examples follow the format shown below.
Example of a custom prompt:
Example of a custom fact extraction prompt:
<CodeGroup>
```python Python
custom_prompt = """
custom_fact_extraction_prompt = """
Please only extract entities containing customer support information, order details, and user information.
Here are some few shot examples:
@@ -67,7 +67,7 @@ Return the facts and customer information in a json format as shown above.
```
</CodeGroup>
Here we initialize the custom prompt in the config:
Here we initialize the custom fact extraction prompt in the config:
<CodeGroup>
```python Python
@@ -82,7 +82,7 @@ config = {
"max_tokens": 2000,
}
},
"custom_prompt": custom_prompt,
"custom_fact_extraction_prompt": custom_fact_extraction_prompt,
"version": "v1.1"
}
@@ -166,4 +166,4 @@ await memory.add('I like going to hikes', { userId: "user123" });
```
</CodeGroup>
The custom prompt will process both the user and assistant messages to extract relevant information according to the defined format.
The custom fact extraction prompt will process both the user and assistant messages to extract relevant information according to the defined format.
@@ -0,0 +1,239 @@
---
title: Custom Update Memory Prompt
icon: "pencil"
iconType: "solid"
---
Update memory prompt is a prompt used to determine the action to be performed on the memory.
By customizing this prompt, you can control how the memory is updated.
## Introduction
Mem0 memory system compares the newly retrieved facts with the existing memory and determines the action to be performed on the memory.
The kinds of actions are:
- Add
- Add the newly retrieved facts to the memory.
- Update
- Update the existing memory with the newly retrieved facts.
- Delete
- Delete the existing memory.
- No Change
- Do not make any changes to the memory.
### Example
Example of a custom update memory prompt:
<CodeGroup>
```python Python
UPDATE_MEMORY_PROMPT = """You are a smart memory manager which controls the memory of a system.
You can perform four operations: (1) add into the memory, (2) update the memory, (3) delete from the memory, and (4) no change.
Based on the above four operations, the memory will change.
Compare newly retrieved facts with the existing memory. For each new fact, decide whether to:
- ADD: Add it to the memory as a new element
- UPDATE: Update an existing memory element
- DELETE: Delete an existing memory element
- NONE: Make no change (if the fact is already present or irrelevant)
There are specific guidelines to select which operation to perform:
1. **Add**: If the retrieved facts contain new information not present in the memory, then you have to add it by generating a new ID in the id field.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "User is a software engineer"
}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "User is a software engineer",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Name is John",
"event" : "ADD"
}
]
}
2. **Update**: If the retrieved facts contain information that is already present in the memory but the information is totally different, then you have to update it.
If the retrieved fact contains information that conveys the same thing as the elements present in the memory, then you have to keep the fact which has the most information.
Example (a) -- if the memory contains "User likes to play cricket" and the retrieved fact is "Loves to play cricket with friends", then update the memory with the retrieved facts.
Example (b) -- if the memory contains "Likes cheese pizza" and the retrieved fact is "Loves cheese pizza", then you do not need to update it because they convey the same information.
If the direction is to update the memory, then you have to update it.
Please keep in mind while updating you have to keep the same ID.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "I really like cheese pizza"
},
{
"id" : "1",
"text" : "User is a software engineer"
},
{
"id" : "2",
"text" : "User likes to play cricket"
}
]
- Retrieved facts: ["Loves chicken pizza", "Loves to play cricket with friends"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Loves cheese and chicken pizza",
"event" : "UPDATE",
"old_memory" : "I really like cheese pizza"
},
{
"id" : "1",
"text" : "User is a software engineer",
"event" : "NONE"
},
{
"id" : "2",
"text" : "Loves to play cricket with friends",
"event" : "UPDATE",
"old_memory" : "User likes to play cricket"
}
]
}
3. **Delete**: If the retrieved facts contain information that contradicts the information present in the memory, then you have to delete it. Or if the direction is to delete the memory, then you have to delete it.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "Name is John"
},
{
"id" : "1",
"text" : "Loves cheese pizza"
}
]
- Retrieved facts: ["Dislikes cheese pizza"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "DELETE"
}
]
}
4. **No Change**: If the retrieved facts contain information that is already present in the memory, then you do not need to make any changes.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "Name is John"
},
{
"id" : "1",
"text" : "Loves cheese pizza"
}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "NONE"
}
]
}
"""
```
</CodeGroup>
## Output format
The prompt needs to guide the output to follow the structure as shown below:
<CodeGroup>
```json Add
{
"memory": [
{
"id" : "0",
"text" : "This information is new",
"event" : "ADD"
}
]
}
```
```json Update
{
"memory": [
{
"id" : "0",
"text" : "This information replaces the old information",
"event" : "UPDATE",
"old_memory" : "Old information"
}
]
}
```
```json Delete
{
"memory": [
{
"id" : "0",
"text" : "This information will be deleted",
"event" : "DELETE"
}
]
}
```
```json No Change
{
"memory": [
{
"id" : "0",
"text" : "No changes for this information",
"event" : "NONE"
}
]
}
```
</CodeGroup>
## custom update memory prompt vs custom prompt
| Feature | `custom_update_memory_prompt` | `custom_prompt` |
|---------|-------------------------------|-----------------|
| Use case | Determine the action to be performed on the memory | Extract the facts from messages |
| Reference | Retrieved facts from messages and old memory | Messages |
| Output | Action to be performed on the memory | Extracted facts |
+62
View File
@@ -0,0 +1,62 @@
---
title: Feedback Mechanism
icon: "thumbs-up"
iconType: "solid"
---
Mem0's **Feedback Mechanism** allows you to provide feedback on the memories generated by your application. This feedback is used to improve the accuracy of the memories and the search results.
## How it works
The feedback mechanism is a simple API that allows you to provide feedback on the memories generated by your application. The feedback is stored in the database and is used to improve the accuracy of the memories and the search results. Over time, Mem0 continuously learns from this feedback, refining its memory generation and search capabilities for better performance.
## Give Feedback
You can give feedback on a memory by calling the `feedback` method on the Mem0 client.
<CodeGroup>
```python Python
from mem0 import Mem0
client = Mem0(api_key="your_api_key")
client.feedback(memory_id="your-memory-id", feedback="NEGATIVE", feedback_reason="I don't like this memory because it is not relevant.")
```
```javascript JavaScript
import MemoryClient from 'mem0ai';
const client = new MemoryClient({ apiKey: 'your-api-key'});
client.feedback({
memory_id: "your-memory-id",
feedback: "NEGATIVE",
feedback_reason: "I don't like this memory because it is not relevant."
})
```
</CodeGroup>
## Feedback Types
The `feedback` parameter can be one of the following values:
- `POSITIVE`: The memory is useful.
- `NEGATIVE`: The memory is not useful.
- `VERY_NEGATIVE`: The memory is not useful at all.
## Parameters
The `feedback` method takes the following parameters:
- `memory_id`: The ID of the memory to give feedback on.
- `feedback`: The feedback to give on the memory. (Optional)
- `feedback_reason`: The reason for the feedback. (Optional)
The `feedback_reason` parameter is optional and can be used to provide a reason for the feedback.
<Note>
You can pass `None` or `null` to the `feedback` and `feedback_reason` parameters to remove the feedback for a memory.
</Note>
+295
View File
@@ -0,0 +1,295 @@
---
title: Graph Memory
icon: "circle-nodes"
iconType: "solid"
description: "Enable graph-based memory retrieval for more contextually relevant results"
---
## Overview
Graph Memory enhances memory pipeline by creating relationships between entities in your data. It builds a network of interconnected information for more contextually relevant search results.
This feature allows your AI applications to understand connections between entities, providing richer context for responses. It's ideal for applications needing relationship tracking and nuanced information retrieval across related memories.
## How Graph Memory Works
The Graph Memory feature analyzes how each entity connects and relates to each other. When enabled:
1. Mem0 automatically builds a graph representation of entities
2. Retrieval considers graph relationships between entities
3. Results include entities that may be contextually important even if they're not direct semantic matches
## Using Graph Memory
To use Graph Memory, you need to enable it in your API calls by setting the `enable_graph=True` parameter. You'll also need to specify `output_format="v1.1"` to receive the enriched response format.
### Adding Memories with Graph Memory
When adding new memories, enable Graph Memory to automatically build relationships with existing memories:
<CodeGroup>
```python Python
from mem0 import MemoryClient
client = MemoryClient(
api_key="your-api-key",
org_id="your-org-id",
project_id="your-project-id"
)
messages = [
{"role": "user", "content": "My name is Joseph"},
{"role": "assistant", "content": "Hello Joseph, it's nice to meet you!"},
{"role": "user", "content": "I'm from Seattle and I work as a software engineer"}
]
# Enable graph memory when adding
client.add(
messages,
user_id="joseph",
version="v1",
enable_graph=True,
output_format="v1.1"
)
```
```javascript JavaScript
import { MemoryClient } from "mem0";
const client = new MemoryClient({
apiKey: "your-api-key",
orgId: "your-org-id",
projectId: "your-project-id"
});
const messages = [
{ role: "user", content: "My name is Joseph" },
{ role: "assistant", content: "Hello Joseph, it's nice to meet you!" },
{ role: "user", content: "I'm from Seattle and I work as a software engineer" }
];
// Enable graph memory when adding
await client.add({
messages,
userId: "joseph",
version: "v1",
enableGraph: true,
outputFormat: "v1.1"
});
```
```json Output
{
"results": [
{
"memory": "Name is Joseph",
"event": "ADD",
"id": "4a5a417a-fa10-43b5-8c53-a77c45e80438"
},
{
"memory": "Is from Seattle",
"event": "ADD",
"id": "8d268d0f-5452-4714-b27d-ae46f676a49d"
},
{
"memory": "Is a software engineer",
"event": "ADD",
"id": "5f0a184e-ddea-4fe6-9b92-692d6a901df8"
}
]
}
```
</CodeGroup>
The graph memory would look like this:
<Frame>
<img src="/images/graph-platform.png" alt="Graph Memory Visualization showing relationships between entities" />
</Frame>
<Caption>Graph Memory creates a network of relationships between entities, enabling more contextual retrieval</Caption>
<Note>
Response for the graph memory's `add` operation will not be available directly in the response.
As adding graph memories is an asynchronous operation due to heavy processing,
you can use the `get_all()` endpoint to retrieve the memory with the graph metadata.
</Note>
### Searching with Graph Memory
When searching memories, Graph Memory helps retrieve entities that are contextually important even if they're not direct semantic matches.
<CodeGroup>
```python Python
# Search with graph memory enabled
results = client.search(
"what is my name?",
user_id="joseph",
enable_graph=True,
output_format="v1.1"
)
print(results)
```
```javascript JavaScript
// Search with graph memory enabled
const results = await client.search({
query: "what is my name?",
userId: "joseph",
enableGraph: true,
outputFormat: "v1.1"
});
console.log(results);
```
```json Output
{
"results": [
{
"id": "4a5a417a-fa10-43b5-8c53-a77c45e80438",
"memory": "Name is Joseph",
"user_id": "joseph",
"metadata": null,
"categories": ["personal_details"],
"immutable": false,
"created_at": "2025-03-19T09:09:00.146390-07:00",
"updated_at": "2025-03-19T09:09:00.146404-07:00",
"score": 0.3621795393335552
},
{
"id": "8d268d0f-5452-4714-b27d-ae46f676a49d",
"memory": "Is from Seattle",
"user_id": "joseph",
"metadata": null,
"categories": ["personal_details"],
"immutable": false,
"created_at": "2025-03-19T09:09:00.170680-07:00",
"updated_at": "2025-03-19T09:09:00.170692-07:00",
"score": 0.31212713194651254
}
],
"relations": [
{
"source": "joseph",
"source_type": "person",
"relationship": "name",
"target": "joseph",
"target_type": "person",
"score": 0.39
}
]
}
```
</CodeGroup>
### Retrieving All Memories with Graph Memory
When retrieving all memories, Graph Memory provides additional relationship context:
<CodeGroup>
```python Python
# Get all memories with graph context
memories = client.get_all(
user_id="joseph",
enable_graph=True,
output_format="v1.1"
)
print(memories)
```
```javascript JavaScript
// Get all memories with graph context
const memories = await client.getAll({
userId: "joseph",
enableGraph: true,
outputFormat: "v1.1"
});
console.log(memories);
```
```json Output
{
"results": [
{
"id": "5f0a184e-ddea-4fe6-9b92-692d6a901df8",
"memory": "Is a software engineer",
"user_id": "joseph",
"metadata": null,
"categories": ["professional_details"],
"immutable": false,
"created_at": "2025-03-19T09:09:00.194116-07:00",
"updated_at": "2025-03-19T09:09:00.194128-07:00",
},
{
"id": "8d268d0f-5452-4714-b27d-ae46f676a49d",
"memory": "Is from Seattle",
"user_id": "joseph",
"metadata": null,
"categories": ["personal_details"],
"immutable": false,
"created_at": "2025-03-19T09:09:00.170680-07:00",
"updated_at": "2025-03-19T09:09:00.170692-07:00",
},
{
"id": "4a5a417a-fa10-43b5-8c53-a77c45e80438",
"memory": "Name is Joseph",
"user_id": "joseph",
"metadata": null,
"categories": ["personal_details"],
"immutable": false,
"created_at": "2025-03-19T09:09:00.146390-07:00",
"updated_at": "2025-03-19T09:09:00.146404-07:00",
}
],
"relations": [
{
"source": "joseph",
"source_type": "person",
"relationship": "name",
"target": "joseph",
"target_type": "person"
},
{
"source": "joseph",
"source_type": "person",
"relationship": "city",
"target": "seattle",
"target_type": "city"
},
{
"source": "joseph",
"source_type": "person",
"relationship": "job",
"target": "software engineer",
"target_type": "job"
}
]
}
```
</CodeGroup>
## Best Practices
- Enable Graph Memory for applications where understanding context and relationships between memories is important
- Graph Memory works best with a rich history of related conversations
- Consider Graph Memory for long-running assistants that need to track evolving information
## Performance Considerations
Graph Memory requires additional processing and may increase response times slightly for very large memory stores. However, for most use cases, the improved retrieval quality outweighs the minimal performance impact.
If you have any questions, please feel free to reach out to us using one of the following methods:
<Snippet file="get-help.mdx" />
+24 -4
View File
@@ -71,13 +71,32 @@ Here's an example schema for extracting professional profile information:
### Submit Export Job
You can optionally provide additional instructions to guide how memories are processed and structured during export using the `export_instructions` parameter.
<CodeGroup>
```python Python
# Basic export request
response = client.create_memory_export(
schema=json_schema,
user_id="user123"
user_id="alice"
)
# Export with custom instructions
export_instructions = """
1. Create a comprehensive profile with detailed information in each category
2. Only mark fields as "None" when absolutely no relevant information exists
3. Base all information directly on the user's memories
4. When contradictions exist, prioritize the most recent information
5. Clearly distinguish between factual statements and inferences
"""
response = client.create_memory_export(
schema=json_schema,
user_id="alice",
export_instructions=export_instructions
)
print(response)
```
@@ -87,7 +106,8 @@ curl -X POST "https://api.mem0.ai/v1/memories/export/" \
-H "Content-Type: application/json" \
-d '{
"schema": {json_schema},
"user_id": "user123"
"user_id": "alice",
"export_instructions": "1. Create a comprehensive profile with detailed information\n2. Only mark fields as \"None\" when absolutely no relevant information exists"
}'
```
@@ -107,12 +127,12 @@ Once the export job is complete, you can retrieve the structured data:
<CodeGroup>
```python Python
response = client.get_memory_export(user_id="user123")
response = client.get_memory_export(user_id="alice")
print(response)
```
```bash cURL
curl -X GET "https://api.mem0.ai/v1/memories/export/?user_id=user123" \
curl -X GET "https://api.mem0.ai/v1/memories/export/?user_id=alice" \
-H "Authorization: Token your-api-key"
```
+6
View File
@@ -12,6 +12,9 @@ Learn about the key features and capabilities that make Mem0 a powerful platform
<Card title="Advanced Retrieval" icon="magnifying-glass" href="/features/advanced-retrieval">
Superior search results using state-of-the-art algorithms, including keyword search, reranking, and filtering capabilities.
</Card>
<Card title="Contextual Add" icon="square-plus" href="/features/contextual-add">
Only send your latest conversation history - we automatically retrieve the rest and generate properly contextualized memories.
</Card>
<Card title="Multimodal Support" icon="photo-film" href="/features/multimodal-support">
Process and analyze various types of content including images.
</Card>
@@ -33,6 +36,9 @@ Learn about the key features and capabilities that make Mem0 a powerful platform
<Card title="Memory Export" icon="file-export" href="/features/memory-export">
Export memories in structured formats using customizable Pydantic schemas.
</Card>
<Card title="Graph Memory" icon="graph" href="/features/graph-memory">
Add memories in the form of nodes and edges in a graph database and search for related memories.
</Card>
</CardGroup>
## Getting Help
Binary file not shown.

After

Width:  |  Height:  |  Size: 72 KiB

+7 -7
View File
@@ -95,7 +95,7 @@ add_input = {
{"role": "user", "content": "Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts."},
{"role": "assistant", "content": "Hello Alex! I've noted that you're a vegetarian and have a nut allergy."}
],
"user_id": "alex123",
"user_id": "alex",
"output_format": "v1.1",
"metadata": {"food": "vegan"}
}
@@ -173,7 +173,7 @@ search_input = {
"filters": {
"AND": [
{"created_at": {"gte": "2024-07-20", "lte": "2024-12-10"}},
{"user_id": "alex123"}
{"user_id": "alex"}
]
},
"version": "v2"
@@ -186,7 +186,7 @@ result = search_tool.invoke(search_input)
{
"id": "1a75e827-7eca-45ea-8c5c-cfd43299f061",
"memory": "Name is Alex",
"user_id": "alex123",
"user_id": "alex",
"hash": "d0fccc8fa47f7a149ee95750c37bb0ca",
"metadata": {
"food": "vegan"
@@ -255,7 +255,7 @@ get_all_input = {
"version": "v2",
"filters": {
"AND": [
{"user_id": "alex123"},
{"user_id": "alex"},
{"created_at": {"gte": "2024-07-01", "lte": "2024-12-31"}}
]
},
@@ -274,7 +274,7 @@ get_all_result = get_all_tool.invoke(get_all_input)
{
"id": "1a75e827-7eca-45ea-8c5c-cfd43299f061",
"memory": "Name is Alex",
"user_id": "alex123",
"user_id": "alex",
"hash": "d0fccc8fa47f7a149ee95750c37bb0ca",
"metadata": {
"food": "vegan"
@@ -288,7 +288,7 @@ get_all_result = get_all_tool.invoke(get_all_input)
{
"id": "91509588-0b39-408a-8df3-84b3bce8c521",
"memory": "Is a vegetarian",
"user_id": "alex123",
"user_id": "alex",
"hash": "ce6b1c84586772ab9995a9477032df99",
"metadata": {
"food": "vegan"
@@ -303,7 +303,7 @@ get_all_result = get_all_tool.invoke(get_all_input)
{
"id": "8d74f7a0-6107-4589-bd6f-210f6bf4fbbb",
"memory": "Is allergic to nuts",
"user_id": "alex123",
"user_id": "alex",
"hash": "7873cd0e5a29c513253d9fad038e758b",
"metadata": {
"food": "vegan"
+14 -2
View File
@@ -22,7 +22,11 @@ pip install mem0ai
<Tabs>
<Tab title="Basic">
```python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "your-api-key"
m = Memory()
```
</Tab>
@@ -42,8 +46,11 @@ docker run -p 6333:6333 -p 6334:6334 \
Then, instantiate memory with qdrant server:
```python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "your-api-key"
config = {
"vector_store": {
"provider": "qdrant",
@@ -61,8 +68,11 @@ m = Memory.from_config(config)
<Tab title="Advanced (Graph Memory)">
```python
import os
from mem0 import Memory
os.environ["OPENAI_API_KEY"] = "your-api-key"
config = {
"graph_store": {
"provider": "neo4j",
@@ -372,7 +382,8 @@ Mem0 offers extensive configuration options to customize its behavior according
|------------------|--------------------------------------|----------------------------|
| `history_db_path` | Path to the history database | "{mem0_dir}/history.db" |
| `version` | API version | "v1.1" |
| `custom_prompt` | Custom prompt for memory processing | None |
| `custom_fact_extraction_prompt` | Custom prompt for memory processing | None |
| `custom_update_memory_prompt` | Custom prompt for update memory | None |
</Accordion>
<Accordion title="Complete Configuration Example">
@@ -409,7 +420,8 @@ config = {
},
"history_db_path": "/path/to/history.db",
"version": "v1.1",
"custom_prompt": "Optional custom prompt for memory processing"
"custom_fact_extraction_prompt": "Optional custom prompt for fact extraction for memory",
"custom_update_memory_prompt": "Optional custom prompt for update memory"
}
```
</Accordion>
+12 -6
View File
@@ -932,27 +932,27 @@
"x-code-samples": [
{
"lang": "Python",
"source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nmessages = [\n {\"role\": \"user\", \"content\": \"<user-message>\"},\n {\"role\": \"assistant\", \"content\": \"<assistant-response>\"}\n]\n\nclient.add(messages, user_id=\"<user-id>\")"
"source": "# To use the Python SDK, install the package:\n# pip install mem0ai\n\nfrom mem0 import MemoryClient\n\nclient = MemoryClient(api_key=\"your_api_key\", org_id=\"your_org_id\", project_id=\"your_project_id\")\n\nmessages = [\n {\"role\": \"user\", \"content\": \"<user-message>\"},\n {\"role\": \"assistant\", \"content\": \"<assistant-response>\"}\n]\n\nclient.add(messages, user_id=\"<user-id>\", version=\"v2\")"
},
{
"lang": "JavaScript",
"source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst messages = [\n { role: \"user\", content: \"Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts.\" },\n { role: \"assistant\", content: \"Hello Alex! I've noted that you're a vegetarian and have a nut allergy. I'll keep this in mind for any food-related recommendations or discussions.\" }\n];\n\nclient.add(messages, { user_id: \"<user_id>\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));"
"source": "// To use the JavaScript SDK, install the package:\n// npm i mem0ai\n\nimport MemoryClient from 'mem0ai';\nconst client = new MemoryClient({ apiKey: \"your-api-key\" });\n\nconst messages = [\n { role: \"user\", content: \"Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts.\" },\n { role: \"assistant\", content: \"Hello Alex! I've noted that you're a vegetarian and have a nut allergy. I'll keep this in mind for any food-related recommendations or discussions.\" }\n];\n\nclient.add(messages, { user_id: \"<user_id>\", version: \"v2\" })\n .then(result => console.log(result))\n .catch(error => console.error(error));"
},
{
"lang": "cURL",
"source": "curl --request POST \\\n --url https://api.mem0.ai/v1/memories/ \\\n --header 'Authorization: Token <api-key>' \\\n --header 'Content-Type: application/json' \\\n --data '{\n \"messages\": [\n {}\n ],\n \"agent_id\": \"<string>\",\n \"user_id\": \"<string>\",\n \"app_id\": \"<string>\",\n \"run_id\": \"<string>\",\n \"metadata\": {},\n \"includes\": \"<string>\",\n \"excludes\": \"<string>\",\n \"infer\": true,\n \"custom_categories\": {}, \n \"org_id\": \"<string>\",\n \"project_id\": \"<string>\"\n}'"
"source": "curl --request POST \\\n --url https://api.mem0.ai/v1/memories/ \\\n --header 'Authorization: Token <api-key>' \\\n --header 'Content-Type: application/json' \\\n --data '{\n \"messages\": [\n {}\n ],\n \"agent_id\": \"<string>\",\n \"user_id\": \"<string>\",\n \"app_id\": \"<string>\",\n \"run_id\": \"<string>\",\n \"metadata\": {},\n \"includes\": \"<string>\",\n \"excludes\": \"<string>\",\n \"infer\": true,\n \"custom_categories\": {}, \n \"org_id\": \"<string>\",\n \"project_id\": \"<string>\",\n \"version\": \"v2\"\n}'"
},
{
"lang": "Go",
"source": "package main\n\nimport (\n\t\"fmt\"\n\t\"strings\"\n\t\"net/http\"\n\t\"io/ioutil\"\n)\n\nfunc main() {\n\n\turl := \"https://api.mem0.ai/v1/memories/\"\n\n\tpayload := strings.NewReader(\"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {},\n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\"\n}\")\n\n\treq, _ := http.NewRequest(\"POST\", url, payload)\n\n\treq.Header.Add(\"Authorization\", \"Token <api-key>\")\n\treq.Header.Add(\"Content-Type\", \"application/json\")\n\n\tres, _ := http.DefaultClient.Do(req)\n\n\tdefer res.Body.Close()\n\tbody, _ := ioutil.ReadAll(res.Body)\n\n\tfmt.Println(res)\n\tfmt.Println(string(body))\n\n}"
"source": "package main\n\nimport (\n\t\"fmt\"\n\t\"strings\"\n\t\"net/http\"\n\t\"io/ioutil\"\n)\n\nfunc main() {\n\n\turl := \"https://api.mem0.ai/v1/memories/\"\n\n\tpayload := strings.NewReader(\"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {},\n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\",\n \\\"version\\\": \"v2\"\n}\")\n\n\treq, _ := http.NewRequest(\"POST\", url, payload)\n\n\treq.Header.Add(\"Authorization\", \"Token <api-key>\")\n\treq.Header.Add(\"Content-Type\", \"application/json\")\n\n\tres, _ := http.DefaultClient.Do(req)\n\n\tdefer res.Body.Close()\n\tbody, _ := ioutil.ReadAll(res.Body)\n\n\tfmt.Println(res)\n\tfmt.Println(string(body))\n\n}"
},
{
"lang": "PHP",
"source": "<?php\n\n$curl = curl_init();\n\ncurl_setopt_array($curl, [\n CURLOPT_URL => \"https://api.mem0.ai/v1/memories/\",\n CURLOPT_RETURNTRANSFER => true,\n CURLOPT_ENCODING => \"\",\n CURLOPT_MAXREDIRS => 10,\n CURLOPT_TIMEOUT => 30,\n CURLOPT_HTTP_VERSION => CURL_HTTP_VERSION_1_1,\n CURLOPT_CUSTOMREQUEST => \"POST\",\n CURLOPT_POSTFIELDS => \"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {}, \n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\"\n}\",\n CURLOPT_HTTPHEADER => [\n \"Authorization: Token <api-key>\",\n \"Content-Type: application/json\"\n ],\n]);\n\n$response = curl_exec($curl);\n$err = curl_error($curl);\n\ncurl_close($curl);\n\nif ($err) {\n echo \"cURL Error #:\" . $err;\n} else {\n echo $response;\n}"
"source": "<?php\n\n$curl = curl_init();\n\ncurl_setopt_array($curl, [\n CURLOPT_URL => \"https://api.mem0.ai/v1/memories/\",\n CURLOPT_RETURNTRANSFER => true,\n CURLOPT_ENCODING => \"\",\n CURLOPT_MAXREDIRS => 10,\n CURLOPT_TIMEOUT => 30,\n CURLOPT_HTTP_VERSION => CURL_HTTP_VERSION_1_1,\n CURLOPT_CUSTOMREQUEST => \"POST\",\n CURLOPT_POSTFIELDS => \"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {}, \n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\",\n \\\"version\\\": \"v2\"\n}\",\n CURLOPT_HTTPHEADER => [\n \"Authorization: Token <api-key>\",\n \"Content-Type: application/json\"\n ],\n]);\n\n$response = curl_exec($curl);\n$err = curl_error($curl);\n\ncurl_close($curl);\n\nif ($err) {\n echo \"cURL Error #:\" . $err;\n} else {\n echo $response;\n}"
},
{
"lang": "Java",
"source": "HttpResponse<String> response = Unirest.post(\"https://api.mem0.ai/v1/memories/\")\n .header(\"Authorization\", \"Token <api-key>\")\n .header(\"Content-Type\", \"application/json\")\n .body(\"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {}, \n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\"\n}\")\n .asString();"
"source": "HttpResponse<String> response = Unirest.post(\"https://api.mem0.ai/v1/memories/\")\n .header(\"Authorization\", \"Token <api-key>\")\n .header(\"Content-Type\", \"application/json\")\n .body(\"{\n \\\"messages\\\": [\n {}\n ],\n \\\"agent_id\\\": \\\"<string>\\\",\n \\\"user_id\\\": \\\"<string>\\\",\n \\\"app_id\\\": \\\"<string>\\\",\n \\\"run_id\\\": \\\"<string>\\\",\n \\\"metadata\\\": {},\n \\\"includes\\\": \\\"<string>\\\",\n \\\"excludes\\\": \\\"<string>\\\",\n \\\"infer\\\": true,\n \\\"custom_categories\\\": {}, \n \\\"org_id\\\": \\\"<string>\\\",\n \\\"project_id\\\": \\\"<string>\",\n \\\"version\\\": \"v2\"\n}\")\n .asString();"
}
],
"x-codegen-request-body-name": "data"
@@ -4809,6 +4809,12 @@
"title": "Project id",
"type": "string",
"nullable": true
},
"version": {
"description": "The version of the memory to use. The default version is v1, which is deprecated. We recommend using v2 for new applications.",
"title": "Version",
"type": "string",
"nullable": true
}
}
},
+45 -41
View File
@@ -87,10 +87,10 @@ messages = [
]
# The default output_format is v1.0
client.add(messages, user_id="alex", output_format="v1.0")
client.add(messages, user_id="alex", output_format="v1.0", version="v2")
# To use the latest output_format, set the output_format parameter to "v1.1"
client.add(messages, user_id="alex", output_format="v1.1", metadata={"food": "vegan"})
client.add(messages, user_id="alex", output_format="v1.1", metadata={"food": "vegan"}, version="v2")
```
```javascript JavaScript
@@ -98,7 +98,7 @@ const messages = [
{"role": "user", "content": "Hi, I'm Alex. I'm a vegetarian and I'm allergic to nuts."},
{"role": "assistant", "content": "Hello Alex! I've noted that you're a vegetarian and have a nut allergy. I'll keep this in mind for any food-related recommendations or discussions."}
];
client.add(messages, { user_id: "alex", output_format: "v1.1", metadata: { food: "vegan" } })
client.add(messages, { user_id: "alex", output_format: "v1.1", metadata: { food: "vegan" }, version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
@@ -116,7 +116,8 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
"output_format": "v1.1",
"metadata": {
"food": "vegan"
}
},
"version": "v2"
}'
```
@@ -191,10 +192,10 @@ messages = [
]
# The default output_format is v1.0
client.add(messages, user_id="alex123", run_id="trip-planning-2024", output_format="v1.0")
client.add(messages, user_id="alex", run_id="trip-planning-2024", output_format="v1.0", version="v2")
# To use the latest output_format, set the output_format parameter to "v1.1"
client.add(messages, user_id="alex123", run_id="trip-planning-2024", output_format="v1.1")
client.add(messages, user_id="alex", run_id="trip-planning-2024", output_format="v1.1", version="v2")
```
```javascript JavaScript
@@ -204,7 +205,7 @@ const messages = [
{"role": "user", "content": "Yes, please! Especially in Tokyo."},
{"role": "assistant", "content": "Great! I'll remember that you're interested in vegetarian restaurants in Tokyo for your upcoming trip. I'll prepare a list for you in our next interaction."}
];
client.add(messages, { user_id: "alex123", run_id: "trip-planning-2024", output_format: "v1.1" })
client.add(messages, { user_id: "alex", run_id: "trip-planning-2024", output_format: "v1.1", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
@@ -220,9 +221,10 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
{"role": "user", "content": "Yes, please! Especially in Tokyo."},
{"role": "assistant", "content": "Great! I'll remember that you're interested in vegetarian restaurants in Tokyo for your upcoming trip. I'll prepare a list for you in our next interaction."}
],
"user_id": "alex123",
"user_id": "alex",
"run_id": "trip-planning-2024",
"output_format": "v1.1"
"output_format": "v1.1",
"version": "v2"
}'
```
@@ -274,10 +276,10 @@ messages = [
]
# The default output_format is v1.0
client.add(messages, agent_id="ai-tutor", output_format="v1.0")
client.add(messages, agent_id="ai-tutor", output_format="v1.0", version="v2")
# To use the latest output_format, set the output_format parameter to "v1.1"
client.add(messages, agent_id="ai-tutor", output_format="v1.1")
client.add(messages, agent_id="ai-tutor", output_format="v1.1", version="v2")
```
```javascript JavaScript
@@ -285,7 +287,7 @@ const messages = [
{"role": "system", "content": "You are an AI tutor with a personality. Give yourself a name for the user."},
{"role": "assistant", "content": "Understood. I'm an AI tutor with a personality. My name is Alice."}
];
client.add(messages, { agent_id: "ai-tutor", output_format: "v1.1" })
client.add(messages, { agent_id: "ai-tutor", output_format: "v1.1", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
@@ -300,7 +302,8 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
{"role": "assistant", "content": "Understood. I'm an AI tutor with a personality. My name is Alice."}
],
"agent_id": "ai-tutor",
"output_format": "v1.1"
"output_format": "v1.1",
"version": "v2"
}'
```
@@ -368,7 +371,7 @@ messages = [
{"role": "assistant", "content": "That's great! I'm going to Dubai next month."},
]
client.add(messages=messages, user_id="user1", agent_id="agent1")
client.add(messages=messages, user_id="user1", agent_id="agent1", version="v2")
```
```javascript JavaScript
@@ -377,7 +380,7 @@ const messages = [
{"role": "assistant", "content": "That's great! I'm going to Dubai next month."},
]
client.add(messages, { user_id: "user1", agent_id: "agent1" })
client.add(messages, { user_id: "user1", agent_id: "agent1", version: "v2" })
.then(response => console.log(response))
.catch(error => console.error(error));
```
@@ -392,7 +395,8 @@ curl -X POST "https://api.mem0.ai/v1/memories/" \
{"role": "assistant", "content": "That's great! I'm going to Dubai next month."},
],
"user_id": "user1",
"agent_id": "agent1"
"agent_id": "agent1",
"version": "v2"
}'
```
@@ -1004,17 +1008,17 @@ curl -X GET "https://api.mem0.ai/v1/memories/?agent_id=ai-tutor&page=1&page_size
<CodeGroup>
```python Python
short_term_memories = client.get_all(user_id="alex123", run_id="trip-planning-2024", page=1, page_size=50)
short_term_memories = client.get_all(user_id="alex", run_id="trip-planning-2024", page=1, page_size=50)
```
```javascript JavaScript
client.getAll({ user_id: "alex123", run_id: "trip-planning-2024", page: 1, page_size: 50 })
client.getAll({ user_id: "alex", run_id: "trip-planning-2024", page: 1, page_size: 50 })
.then(memories => console.log(memories))
.catch(error => console.error(error));
```
```bash cURL
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planning-2024&page=1&page_size=50" \
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&run_id=trip-planning-2024&page=1&page_size=50" \
-H "Authorization: Token your-api-key"
```
@@ -1027,7 +1031,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id":"06d8df63-7bd2-4fad-9acb-60871bcecee0",
"memory":"Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id":"alex123",
"user_id":"alex",
"hash":"d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1038,7 +1042,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id":"b4229775-d860-4ccb-983f-0f628ca112f5",
"memory":"Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id":"alex123",
"user_id":"alex",
"hash":"d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1049,7 +1053,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id":"df1aca24-76cf-4b92-9f58-d03857efcb64",
"memory":"Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id":"alex123",
"user_id":"alex",
"hash":"d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1072,7 +1076,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id": "06d8df63-7bd2-4fad-9acb-60871bcecee0",
"memory": "Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id": "alex123",
"user_id": "alex",
"hash": "d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1083,7 +1087,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id": "b4229775-d860-4ccb-983f-0f628ca112f5",
"memory": "Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id": "alex123",
"user_id": "alex",
"hash": "d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1094,7 +1098,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&run_id=trip-planni
{
"id": "df1aca24-76cf-4b92-9f58-d03857efcb64",
"memory": "Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id": "alex123",
"user_id": "alex",
"hash": "d2088c936e259f2f5d2d75543d31401c",
"metadata":None,
"immutable": false,
@@ -1133,7 +1137,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/582bbe6d-506b-48c6-a4c6-5df3b1e6342
{
"id":"06d8df63-7bd2-4fad-9acb-60871bcecee0",
"memory":"Planning a trip to Japan next month. Interested in vegetarian restaurants in Tokyo.",
"user_id":"alex123",
"user_id":"alex",
"hash":"d2088c936e259f2f5d2d75543d31401c",
"metadata":"None",
"immutable": false,
@@ -1152,55 +1156,55 @@ You can filter memories by their categories when using get_all:
```python Python
# Get memories with specific categories
memories = client.get_all(user_id="alex123", categories=["likes"])
memories = client.get_all(user_id="alex", categories=["likes"])
# Get memories with multiple categories
memories = client.get_all(user_id="alex123", categories=["likes", "food_preferences"])
memories = client.get_all(user_id="alex", categories=["likes", "food_preferences"])
# Custom pagination with categories
memories = client.get_all(user_id="alex123", categories=["likes"], page=1, page_size=50)
memories = client.get_all(user_id="alex", categories=["likes"], page=1, page_size=50)
# Get memories with specific keywords
memories = client.get_all(user_id="alex123", keywords="to play", page=1, page_size=50)
memories = client.get_all(user_id="alex", keywords="to play", page=1, page_size=50)
```
```javascript JavaScript
// Get memories with specific categories
client.getAll({ user_id: "alex123", categories: ["likes"] })
client.getAll({ user_id: "alex", categories: ["likes"] })
.then(memories => console.log(memories))
.catch(error => console.error(error));
// Get memories with multiple categories
client.getAll({ user_id: "alex123", categories: ["likes", "food_preferences"] })
client.getAll({ user_id: "alex", categories: ["likes", "food_preferences"] })
.then(memories => console.log(memories))
.catch(error => console.error(error));
// Custom pagination with categories
client.getAll({ user_id: "alex123", categories: ["likes"], page: 1, page_size: 50 })
client.getAll({ user_id: "alex", categories: ["likes"], page: 1, page_size: 50 })
.then(memories => console.log(memories))
.catch(error => console.error(error));
// Get memories with specific keywords
client.getAll({ user_id: "alex123", keywords: "to play", page: 1, page_size: 50 })
client.getAll({ user_id: "alex", keywords: "to play", page: 1, page_size: 50 })
.then(memories => console.log(memories))
.catch(error => console.error(error));
```
```bash cURL
# Get memories with specific categories
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&categories=likes" \
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&categories=likes" \
-H "Authorization: Token your-api-key"
# Get memories with multiple categories
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&categories=likes,food_preferences" \
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&categories=likes,food_preferences" \
-H "Authorization: Token your-api-key"
# Custom pagination with categories
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&categories=likes&page=1&page_size=50" \
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&categories=likes&page=1&page_size=50" \
-H "Authorization: Token your-api-key"
# Get memories with specific keywords
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&keywords=to play&page=1&page_size=50" \
curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex&keywords=to play&page=1&page_size=50" \
-H "Authorization: Token your-api-key"
```
@@ -1213,7 +1217,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&keywords=to play&p
{
"id": "06d8df63-7bd2-4fad-9acb-60871bcecee0",
"memory": "Likes pizza and pasta",
"user_id": "alex123",
"user_id": "alex",
"hash": "d2088c936e259f2f5d2d75543d31401c",
"metadata": null,
"immutable": false,
@@ -1224,7 +1228,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/?user_id=alex123&keywords=to play&p
{
"id": "b4229775-d860-4ccb-983f-0f628ca112f5",
"memory": "Likes to travel to beach destinations",
"user_id": "alex123",
"user_id": "alex",
"hash": "d2088c936e259f2f5d2d75543d31401c",
"metadata": null,
"immutable": false,
@@ -1585,7 +1589,7 @@ curl -X GET "https://api.mem0.ai/v1/memories/<memory-id-here>/history/" \
],
"old_memory":"None",
"new_memory":"Turned vegetarian.",
"user_id":"alex123456",
"user_id":"alex",
"event":"ADD",
"metadata":"None",
"created_at":"2024-07-26T01:02:41.737310-07:00",
+9 -5
View File
@@ -1,6 +1,6 @@
{
"name": "mem0ai",
"version": "2.1.4",
"version": "2.1.8",
"description": "The Memory Layer For Your AI Apps",
"main": "./dist/index.js",
"module": "./dist/index.mjs",
@@ -31,9 +31,10 @@
"dist"
],
"scripts": {
"clean": "rm -rf dist",
"build": "npm run clean && prettier --check . && tsup",
"dev": "nodemon",
"clean": "rimraf dist",
"build": "npm run clean && npx prettier --check . && npx tsup",
"dev": "npx nodemon",
"start": "npx ts-node src/oss/examples/basic.ts",
"test": "jest",
"test:ts": "jest --config jest.config.js",
"test:watch": "jest --config jest.config.js --watch",
@@ -74,7 +75,9 @@
"dotenv": "^16.4.5",
"fix-tsup-cjs": "^1.2.0",
"jest": "^29.7.0",
"nodemon": "^3.0.1",
"prettier": "^3.5.2",
"rimraf": "^5.0.5",
"ts-jest": "^29.2.6",
"ts-node": "^10.9.2",
"tsup": "^8.3.0",
@@ -96,7 +99,8 @@
"groq-sdk": "0.3.0",
"pg": "8.11.3",
"redis": "4.7.0",
"sqlite3": "5.1.7"
"sqlite3": "5.1.7",
"ollama": "^0.5.14"
},
"peerDependenciesMeta": {
"posthog-node": {
+2
View File
@@ -23,6 +23,8 @@ export type {
Message,
AllUsers,
User,
FeedbackPayload,
Feedback,
} from "./mem0.types";
// Export telemetry types
+13
View File
@@ -12,6 +12,7 @@ import {
Webhook,
WebhookPayload,
Message,
FeedbackPayload,
} from "./mem0.types";
import { captureClientEvent, generateHash } from "./telemetry";
@@ -560,6 +561,18 @@ export default class MemoryClient {
);
return response;
}
async feedback(data: FeedbackPayload): Promise<{ message: string }> {
const response = await this._fetchWithErrorHandling(
`${this.host}/v1/feedback/`,
{
method: "POST",
headers: this.headers,
body: JSON.stringify(data),
},
);
return response;
}
}
export { MemoryClient };
+12
View File
@@ -31,6 +31,12 @@ export enum API_VERSION {
V2 = "v2",
}
export enum Feedback {
POSITIVE = "POSITIVE",
NEGATIVE = "NEGATIVE",
VERY_NEGATIVE = "VERY_NEGATIVE",
}
export interface MultiModalMessages {
type: "image_url";
image_url: {
@@ -164,3 +170,9 @@ export interface WebhookPayload {
name: string;
url: string;
}
export interface FeedbackPayload {
memory_id: string;
feedback?: Feedback | null;
feedback_reason?: string | null;
}
+34 -1
View File
@@ -116,6 +116,36 @@ async function runTests(memory: Memory) {
}
}
async function demoLocalMemory() {
console.log("\n=== Testing In-Memory Vector Store with Ollama===\n");
const memory = new Memory({
version: "v1.1",
embedder: {
provider: "ollama",
config: {
model: "nomic-embed-text:latest",
},
},
vectorStore: {
provider: "memory",
config: {
collectionName: "memories",
dimension: 768, // 768 is the dimension of the nomic-embed-text model
},
},
llm: {
provider: "ollama",
config: {
model: "llama3.1:8b",
},
},
// historyDbPath: "memory.db",
});
await runTests(memory);
}
async function demoMemoryStore() {
console.log("\n=== Testing In-Memory Vector Store ===\n");
@@ -346,6 +376,9 @@ async function main() {
// Test in-memory store
await demoMemoryStore();
// Test in-memory store with Ollama
await demoLocalMemory();
// Test graph memory if Neo4j environment variables are set
if (
process.env.NEO4J_URL &&
@@ -384,4 +417,4 @@ async function main() {
}
}
// main();
main();
+93
View File
@@ -0,0 +1,93 @@
import { Memory } from "../src";
import { Ollama } from "ollama";
import * as readline from "readline";
const memory = new Memory({
embedder: {
provider: "ollama",
config: {
model: "nomic-embed-text:latest",
},
},
vectorStore: {
provider: "memory",
config: {
collectionName: "memories",
dimension: 768, // since we are using nomic-embed-text
},
},
llm: {
provider: "ollama",
config: {
model: "llama3.1:8b",
},
},
historyDbPath: "local-llms.db",
});
async function chatWithMemories(message: string, userId = "default_user") {
const relevantMemories = await memory.search(message, { userId: userId });
const memoriesStr = relevantMemories.results
.map((entry) => `- ${entry.memory}`)
.join("\n");
const systemPrompt = `You are a helpful AI. Answer the question based on query and memories.
User Memories:
${memoriesStr}`;
const messages = [
{ role: "system", content: systemPrompt },
{ role: "user", content: message },
];
const ollama = new Ollama();
const response = await ollama.chat({
model: "llama3.1:8b",
messages: messages,
});
const assistantResponse = response.message.content || "";
messages.push({ role: "assistant", content: assistantResponse });
await memory.add(messages, { userId: userId });
return assistantResponse;
}
async function main() {
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout,
});
console.log("Chat with AI (type 'exit' to quit)");
const askQuestion = (): Promise<string> => {
return new Promise((resolve) => {
rl.question("You: ", (input) => {
resolve(input.trim());
});
});
};
try {
while (true) {
const userInput = await askQuestion();
if (userInput.toLowerCase() === "exit") {
console.log("Goodbye!");
rl.close();
break;
}
const response = await chatWithMemories(userInput, "sample_user");
console.log(`AI: ${response}`);
}
} catch (error) {
console.error("An error occurred:", error);
rl.close();
}
}
main().catch(console.error);
+52
View File
@@ -0,0 +1,52 @@
import { Ollama } from "ollama";
import { Embedder } from "./base";
import { EmbeddingConfig } from "../types";
import { logger } from "../utils/logger";
export class OllamaEmbedder implements Embedder {
private ollama: Ollama;
private model: string;
// Using this variable to avoid calling the Ollama server multiple times
private initialized: boolean = false;
constructor(config: EmbeddingConfig) {
this.ollama = new Ollama({
host: config.url || "http://localhost:11434",
});
this.model = config.model || "nomic-embed-text:latest";
this.ensureModelExists().catch((err) => {
logger.error(`Error ensuring model exists: ${err}`);
});
}
async embed(text: string): Promise<number[]> {
try {
await this.ensureModelExists();
} catch (err) {
logger.error(`Error ensuring model exists: ${err}`);
}
const response = await this.ollama.embeddings({
model: this.model,
prompt: text,
});
return response.embedding;
}
async embedBatch(texts: string[]): Promise<number[][]> {
const response = await Promise.all(texts.map((text) => this.embed(text)));
return response;
}
private async ensureModelExists(): Promise<boolean> {
if (this.initialized) {
return true;
}
const local_models = await this.ollama.list();
if (!local_models.models.find((m: any) => m.name === this.model)) {
logger.info(`Pulling model ${this.model}...`);
await this.ollama.pull({ model: this.model });
}
this.initialized = true;
return true;
}
}
+104
View File
@@ -0,0 +1,104 @@
import { Ollama } from "ollama";
import { LLM, LLMResponse } from "./base";
import { LLMConfig, Message } from "../types";
import { logger } from "../utils/logger";
export class OllamaLLM implements LLM {
private ollama: Ollama;
private model: string;
// Using this variable to avoid calling the Ollama server multiple times
private initialized: boolean = false;
constructor(config: LLMConfig) {
this.ollama = new Ollama({
host: config.config?.url || "http://localhost:11434",
});
this.model = config.model || "llama3.1:8b";
this.ensureModelExists().catch((err) => {
logger.error(`Error ensuring model exists: ${err}`);
});
}
async generateResponse(
messages: Message[],
responseFormat?: { type: string },
tools?: any[],
): Promise<string | LLMResponse> {
try {
await this.ensureModelExists();
} catch (err) {
logger.error(`Error ensuring model exists: ${err}`);
}
const completion = await this.ollama.chat({
model: this.model,
messages: messages.map((msg) => {
const role = msg.role as "system" | "user" | "assistant";
return {
role,
content:
typeof msg.content === "string"
? msg.content
: JSON.stringify(msg.content),
};
}),
...(responseFormat?.type === "json_object" && { format: "json" }),
...(tools && { tools, tool_choice: "auto" }),
});
const response = completion.message;
if (response.tool_calls) {
return {
content: response.content || "",
role: response.role,
toolCalls: response.tool_calls.map((call) => ({
name: call.function.name,
arguments: JSON.stringify(call.function.arguments),
})),
};
}
return response.content || "";
}
async generateChat(messages: Message[]): Promise<LLMResponse> {
try {
await this.ensureModelExists();
} catch (err) {
logger.error(`Error ensuring model exists: ${err}`);
}
const completion = await this.ollama.chat({
messages: messages.map((msg) => {
const role = msg.role as "system" | "user" | "assistant";
return {
role,
content:
typeof msg.content === "string"
? msg.content
: JSON.stringify(msg.content),
};
}),
model: this.model,
});
const response = completion.message;
return {
content: response.content || "",
role: response.role,
};
}
private async ensureModelExists(): Promise<boolean> {
if (this.initialized) {
return true;
}
const local_models = await this.ollama.list();
if (!local_models.models.find((m: any) => m.name === this.model)) {
logger.info(`Pulling model ${this.model}...`);
await this.ollama.pull({ model: this.model });
}
this.initialized = true;
return true;
}
}
+2 -1
View File
@@ -13,8 +13,9 @@ export interface Message {
}
export interface EmbeddingConfig {
apiKey: string;
apiKey?: string;
model?: string;
url?: string;
}
export interface VectorStoreConfig {
+6 -1
View File
@@ -1,4 +1,5 @@
import { OpenAIEmbedder } from "../embeddings/openai";
import { OllamaEmbedder } from "../embeddings/ollama";
import { OpenAILLM } from "../llms/openai";
import { OpenAIStructuredLLM } from "../llms/openai_structured";
import { AnthropicLLM } from "../llms/anthropic";
@@ -10,12 +11,14 @@ import { LLM } from "../llms/base";
import { VectorStore } from "../vector_stores/base";
import { Qdrant } from "../vector_stores/qdrant";
import { RedisDB } from "../vector_stores/redis";
import { OllamaLLM } from "../llms/ollama";
export class EmbedderFactory {
static create(provider: string, config: EmbeddingConfig): Embedder {
switch (provider.toLowerCase()) {
case "openai":
return new OpenAIEmbedder(config);
case "ollama":
return new OllamaEmbedder(config);
default:
throw new Error(`Unsupported embedder provider: ${provider}`);
}
@@ -33,6 +36,8 @@ export class LLMFactory {
return new AnthropicLLM(config);
case "groq":
return new GroqLLM(config);
case "ollama":
return new OllamaLLM(config);
default:
throw new Error(`Unsupported LLM provider: ${provider}`);
}
+48 -20
View File
@@ -62,7 +62,7 @@ interface RedisDocument {
run_id?: string;
user_id?: string;
metadata?: string;
vector_score?: number;
__vector_score?: number;
};
}
@@ -108,6 +108,30 @@ const EXCLUDED_KEYS = new Set([
"updated_at",
]);
// Utility function to convert object keys to snake_case
function toSnakeCase(obj: Record<string, any>): Record<string, any> {
if (typeof obj !== "object" || obj === null) return obj;
return Object.fromEntries(
Object.entries(obj).map(([key, value]) => [
key.replace(/[A-Z]/g, (letter) => `_${letter.toLowerCase()}`),
value,
]),
);
}
// Utility function to convert object keys to camelCase
function toCamelCase(obj: Record<string, any>): Record<string, any> {
if (typeof obj !== "object" || obj === null) return obj;
return Object.fromEntries(
Object.entries(obj).map(([key, value]) => [
key.replace(/_([a-z])/g, (_, letter) => letter.toUpperCase()),
value,
]),
);
}
export class RedisDB implements VectorStore {
private client: RedisClientType<
RedisDefaultModules & RedisModules & RedisFunctions & RedisScripts
@@ -272,7 +296,7 @@ export class RedisDB implements VectorStore {
payloads: Record<string, any>[],
): Promise<void> {
const data = vectors.map((vector, idx) => {
const payload = payloads[idx];
const payload = toSnakeCase(payloads[idx]);
const id = ids[idx];
// Create entry with required fields
@@ -322,8 +346,9 @@ export class RedisDB implements VectorStore {
limit: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
const filterExpr = filters
? Object.entries(filters)
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
const filterExpr = snakeFilters
? Object.entries(snakeFilters)
.filter(([_, value]) => value !== null)
.map(([key, value]) => `@${key}:{${value}}`)
.join(" ")
@@ -344,8 +369,9 @@ export class RedisDB implements VectorStore {
"memory",
"metadata",
"created_at",
"__vector_score",
],
SORTBY: "vector_score",
SORTBY: "__vector_score",
DIALECT: 2,
LIMIT: {
from: 0,
@@ -356,12 +382,12 @@ export class RedisDB implements VectorStore {
try {
const results = (await this.client.ft.search(
this.indexName,
`${filterExpr} =>[KNN ${limit} @embedding $vec AS vector_score]`,
`${filterExpr} =>[KNN ${limit} @embedding $vec AS __vector_score]`,
searchOptions,
)) as unknown as RedisSearchResult;
return results.documents.map((doc) => {
const payload = {
const resultPayload = {
hash: doc.value.hash,
data: doc.value.memory,
created_at: new Date(parseInt(doc.value.created_at)).toISOString(),
@@ -376,8 +402,8 @@ export class RedisDB implements VectorStore {
return {
id: doc.value.memory_id,
payload,
score: doc.value.vector_score,
payload: toCamelCase(resultPayload),
score: Number(doc.value.__vector_score) ?? 0,
};
});
} catch (error) {
@@ -493,26 +519,27 @@ export class RedisDB implements VectorStore {
vector: number[],
payload: Record<string, any>,
): Promise<void> {
const snakePayload = toSnakeCase(payload);
const entry: Record<string, any> = {
memory_id: vectorId,
hash: payload.hash,
memory: payload.data,
created_at: new Date(payload.created_at).getTime(),
updated_at: new Date(payload.updated_at).getTime(),
hash: snakePayload.hash,
memory: snakePayload.data,
created_at: new Date(snakePayload.created_at).getTime(),
updated_at: new Date(snakePayload.updated_at).getTime(),
embedding: Buffer.from(new Float32Array(vector).buffer),
};
// Add optional fields
["agent_id", "run_id", "user_id"].forEach((field) => {
if (field in payload) {
entry[field] = payload[field];
if (field in snakePayload) {
entry[field] = snakePayload[field];
}
});
// Add metadata excluding specific keys
entry.metadata = JSON.stringify(
Object.fromEntries(
Object.entries(payload).filter(([key]) => !EXCLUDED_KEYS.has(key)),
Object.entries(snakePayload).filter(([key]) => !EXCLUDED_KEYS.has(key)),
),
);
@@ -557,8 +584,9 @@ export class RedisDB implements VectorStore {
filters?: SearchFilters,
limit: number = 100,
): Promise<[VectorStoreResult[], number]> {
const filterExpr = filters
? Object.entries(filters)
const snakeFilters = filters ? toSnakeCase(filters) : undefined;
const filterExpr = snakeFilters
? Object.entries(snakeFilters)
.filter(([_, value]) => value !== null)
.map(([key, value]) => `@${key}:{${value}}`)
.join(" ")
@@ -581,7 +609,7 @@ export class RedisDB implements VectorStore {
const items = results.documents.map((doc) => ({
id: doc.value.memory_id,
payload: {
payload: toCamelCase({
hash: doc.value.hash,
data: doc.value.memory,
created_at: new Date(parseInt(doc.value.created_at)).toISOString(),
@@ -592,7 +620,7 @@ export class RedisDB implements VectorStore {
...(doc.value.run_id && { run_id: doc.value.run_id }),
...(doc.value.user_id && { user_id: doc.value.user_id }),
...JSON.parse(doc.value.metadata || "{}"),
},
}),
}));
return [items, results.total];
+20
View File
@@ -3,6 +3,7 @@ import os
import warnings
from functools import wraps
from typing import Any, Dict, List, Optional, Union
from enum import Enum
import httpx
@@ -998,3 +999,22 @@ class AsyncMemoryClient:
response.raise_for_status()
capture_client_event("async_client.delete_webhook", self.sync_client, {"webhook_id": webhook_id})
return response.json()
@api_error_handler
async def feedback(self, memory_id: str, feedback: Optional[str] = None, feedback_reason: Optional[str] = None) -> Dict[str, str]:
VALID_FEEDBACK_VALUES = {"POSITIVE", "NEGATIVE", "VERY_NEGATIVE"}
feedback = feedback.upper() if feedback else None
if feedback is not None and feedback not in VALID_FEEDBACK_VALUES:
raise ValueError(f'feedback must be one of {", ".join(VALID_FEEDBACK_VALUES)} or None')
data = {
"memory_id": memory_id,
"feedback": feedback,
"feedback_reason": feedback_reason
}
response = await self.async_client.post("/v1/feedback/", json=data)
response.raise_for_status()
capture_client_event("async_client.feedback", self.sync_client, data)
return response.json()
+6 -2
View File
@@ -48,8 +48,12 @@ class MemoryConfig(BaseModel):
description="The version of the API",
default="v1.1",
)
custom_prompt: Optional[str] = Field(
description="Custom prompt for the memory",
custom_fact_extraction_prompt: Optional[str] = Field(
description="Custom prompt for the fact extraction",
default=None,
)
custom_update_memory_prompt: Optional[str] = Field(
description="Custom prompt for the update memory",
default=None,
)
+165 -145
View File
@@ -58,169 +58,189 @@ Following is a conversation between the user and the assistant. You have to extr
You should detect the language of the user input and record the facts in the same language.
"""
DEFAULT_UPDATE_MEMORY_PROMPT = """You are a smart memory manager which controls the memory of a system.
You can perform four operations: (1) add into the memory, (2) update the memory, (3) delete from the memory, and (4) no change.
def get_update_memory_messages(retrieved_old_memory_dict, response_content):
return f"""You are a smart memory manager which controls the memory of a system.
You can perform four operations: (1) add into the memory, (2) update the memory, (3) delete from the memory, and (4) no change.
Based on the above four operations, the memory will change.
Based on the above four operations, the memory will change.
Compare newly retrieved facts with the existing memory. For each new fact, decide whether to:
- ADD: Add it to the memory as a new element
- UPDATE: Update an existing memory element
- DELETE: Delete an existing memory element
- NONE: Make no change (if the fact is already present or irrelevant)
Compare newly retrieved facts with the existing memory. For each new fact, decide whether to:
- ADD: Add it to the memory as a new element
- UPDATE: Update an existing memory element
- DELETE: Delete an existing memory element
- NONE: Make no change (if the fact is already present or irrelevant)
There are specific guidelines to select which operation to perform:
There are specific guidelines to select which operation to perform:
1. **Add**: If the retrieved facts contain new information not present in the memory, then you have to add it by generating a new ID in the id field.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "User is a software engineer"
}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "User is a software engineer",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Name is John",
"event" : "ADD"
}
]
1. **Add**: If the retrieved facts contain new information not present in the memory, then you have to add it by generating a new ID in the id field.
- **Example**:
- Old Memory:
[
{{
"id" : "0",
"text" : "User is a software engineer"
}}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{{
"memory" : [
{{
"id" : "0",
"text" : "User is a software engineer",
"event" : "NONE"
}},
{{
"id" : "1",
"text" : "Name is John",
"event" : "ADD"
}}
]
}
}}
2. **Update**: If the retrieved facts contain information that is already present in the memory but the information is totally different, then you have to update it.
If the retrieved fact contains information that conveys the same thing as the elements present in the memory, then you have to keep the fact which has the most information.
Example (a) -- if the memory contains "User likes to play cricket" and the retrieved fact is "Loves to play cricket with friends", then update the memory with the retrieved facts.
Example (b) -- if the memory contains "Likes cheese pizza" and the retrieved fact is "Loves cheese pizza", then you do not need to update it because they convey the same information.
If the direction is to update the memory, then you have to update it.
Please keep in mind while updating you have to keep the same ID.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{{
"id" : "0",
"text" : "I really like cheese pizza"
}},
{{
"id" : "1",
"text" : "User is a software engineer"
}},
{{
"id" : "2",
"text" : "User likes to play cricket"
}}
]
- Retrieved facts: ["Loves chicken pizza", "Loves to play cricket with friends"]
- New Memory:
{{
"memory" : [
{{
"id" : "0",
"text" : "Loves cheese and chicken pizza",
"event" : "UPDATE",
"old_memory" : "I really like cheese pizza"
}},
{{
"id" : "1",
"text" : "User is a software engineer",
"event" : "NONE"
}},
{{
"id" : "2",
"text" : "Loves to play cricket with friends",
"event" : "UPDATE",
"old_memory" : "User likes to play cricket"
}}
]
}}
2. **Update**: If the retrieved facts contain information that is already present in the memory but the information is totally different, then you have to update it.
If the retrieved fact contains information that conveys the same thing as the elements present in the memory, then you have to keep the fact which has the most information.
Example (a) -- if the memory contains "User likes to play cricket" and the retrieved fact is "Loves to play cricket with friends", then update the memory with the retrieved facts.
Example (b) -- if the memory contains "Likes cheese pizza" and the retrieved fact is "Loves cheese pizza", then you do not need to update it because they convey the same information.
If the direction is to update the memory, then you have to update it.
Please keep in mind while updating you have to keep the same ID.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "I really like cheese pizza"
},
{
"id" : "1",
"text" : "User is a software engineer"
},
{
"id" : "2",
"text" : "User likes to play cricket"
}
]
- Retrieved facts: ["Loves chicken pizza", "Loves to play cricket with friends"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Loves cheese and chicken pizza",
"event" : "UPDATE",
"old_memory" : "I really like cheese pizza"
},
{
"id" : "1",
"text" : "User is a software engineer",
"event" : "NONE"
},
{
"id" : "2",
"text" : "Loves to play cricket with friends",
"event" : "UPDATE",
"old_memory" : "User likes to play cricket"
}
]
}
3. **Delete**: If the retrieved facts contain information that contradicts the information present in the memory, then you have to delete it. Or if the direction is to delete the memory, then you have to delete it.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{{
"id" : "0",
"text" : "Name is John"
}},
{{
"id" : "1",
"text" : "Loves cheese pizza"
}}
]
- Retrieved facts: ["Dislikes cheese pizza"]
- New Memory:
{{
"memory" : [
{{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
}},
{{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "DELETE"
}}
]
}}
3. **Delete**: If the retrieved facts contain information that contradicts the information present in the memory, then you have to delete it. Or if the direction is to delete the memory, then you have to delete it.
Please note to return the IDs in the output from the input IDs only and do not generate any new ID.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "Name is John"
},
{
"id" : "1",
"text" : "Loves cheese pizza"
}
]
- Retrieved facts: ["Dislikes cheese pizza"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "DELETE"
}
]
}
4. **No Change**: If the retrieved facts contain information that is already present in the memory, then you do not need to make any changes.
- **Example**:
- Old Memory:
[
{{
"id" : "0",
"text" : "Name is John"
}},
{{
"id" : "1",
"text" : "Loves cheese pizza"
}}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{{
"memory" : [
{{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
}},
{{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "NONE"
}}
]
}}
4. **No Change**: If the retrieved facts contain information that is already present in the memory, then you do not need to make any changes.
- **Example**:
- Old Memory:
[
{
"id" : "0",
"text" : "Name is John"
},
{
"id" : "1",
"text" : "Loves cheese pizza"
}
]
- Retrieved facts: ["Name is John"]
- New Memory:
{
"memory" : [
{
"id" : "0",
"text" : "Name is John",
"event" : "NONE"
},
{
"id" : "1",
"text" : "Loves cheese pizza",
"event" : "NONE"
}
]
}
"""
def get_update_memory_messages(retrieved_old_memory_dict, response_content, custom_update_memory_prompt=None):
if custom_update_memory_prompt is None:
global DEFAULT_UPDATE_MEMORY_PROMPT
custom_update_memory_prompt = DEFAULT_UPDATE_MEMORY_PROMPT
return f"""{custom_update_memory_prompt}
Below is the current content of my memory which I have collected till now. You have to update it in the following format only:
``
```
{retrieved_old_memory_dict}
``
```
The new retrieved facts are mentioned in the triple backticks. You have to analyze the new retrieved facts and determine whether these facts should be added, updated, or deleted in the memory.
```
{response_content}
```
You must return your response in the following JSON structure only:
{{
"memory" : [
{{
"id" : "<ID of the memory>", # Use existing ID for updates/deletes, or new ID for additions
"text" : "<Content of the memory>", # Content of the memory
"event" : "<Operation to be performed>", # Must be "ADD", "UPDATE", "DELETE", or "NONE"
"old_memory" : "<Old memory content>" # Required only if the event is "UPDATE"
}},
...
]
}}
Follow the instruction mentioned below:
- Do not return anything from the custom few shot prompts provided above.
- If the current memory is empty, then you have to add the new retrieved facts to the memory.
@@ -230,4 +250,4 @@ def get_update_memory_messages(retrieved_old_memory_dict, response_content):
- If there is an update, the ID key should remain the same and only the value needs to be updated.
Do not return anything except the JSON format.
"""
"""
+16 -11
View File
@@ -8,21 +8,26 @@ class AzureAISearchConfig(BaseModel):
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")
compression_type: Optional[str] = Field(
None,
description="Type of vector compression to use. Options: 'scalar', 'binary', or None"
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)"
False,
description="Whether to store vectors in half precision (Edm.Half) instead of full precision (Edm.Single)",
)
hybrid_search: bool = Field(
False, description="Whether to use hybrid search. If True, vector_filter_mode must be 'preFilter'"
)
vector_filter_mode: Optional[str] = Field(
"preFilter", description="Mode for vector filtering. Options: 'preFilter', 'postFilter'"
)
@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(
@@ -30,13 +35,13 @@ class AzureAISearchConfig(BaseModel):
"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)}. "
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"]
@@ -45,9 +50,9 @@ class AzureAISearchConfig(BaseModel):
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,
}
}
+5 -1
View File
@@ -1,4 +1,5 @@
from typing import Any, Dict, Optional
from collections.abc import Callable
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field, model_validator
@@ -15,6 +16,9 @@ class ElasticsearchConfig(BaseModel):
verify_certs: bool = Field(True, description="Verify SSL certificates")
use_ssl: bool = Field(True, description="Use SSL for connection")
auto_create_index: bool = Field(True, description="Automatically create index during initialization")
custom_search_query: Optional[Callable[[List[float], int, Optional[Dict]], Dict]] = Field(
None, description="Custom search query function. Parameters: (query, limit, filters) -> Dict"
)
@model_validator(mode="before")
@classmethod
+2 -1
View File
@@ -14,6 +14,7 @@ class OpenSearchConfig(BaseModel):
verify_certs: bool = Field(False, description="Verify SSL certificates (default False for OpenSearch)")
use_ssl: bool = Field(False, description="Use SSL for connection (default False for OpenSearch)")
auto_create_index: bool = Field(True, description="Automatically create index during initialization")
http_auth: Optional[object] = Field(None, description="HTTP authentication method / AWS SigV4")
@model_validator(mode="before")
@classmethod
@@ -23,7 +24,7 @@ class OpenSearchConfig(BaseModel):
raise ValueError("Host must be provided for OpenSearch")
# Authentication: Either API key or user/password must be provided
if not any([values.get("api_key"), (values.get("user") and values.get("password"))]):
if not any([values.get("api_key"), (values.get("user") and values.get("password")), values.get("http_auth")]):
raise ValueError("Either api_key or user/password must be provided for OpenSearch authentication")
return values
+56
View File
@@ -0,0 +1,56 @@
import os
from typing import Any, Dict, Optional
from pydantic import BaseModel, Field, model_validator
class PineconeConfig(BaseModel):
"""Configuration for Pinecone vector database."""
collection_name: str = Field("mem0", description="Name of the index/collection")
embedding_model_dims: int = Field(1536, description="Dimensions of the embedding model")
client: Optional[Any] = Field(None, description="Existing Pinecone client instance")
api_key: Optional[str] = Field(None, description="API key for Pinecone")
environment: Optional[str] = Field(None, description="Pinecone environment")
serverless_config: Optional[Dict[str, Any]] = Field(None, description="Configuration for serverless deployment")
pod_config: Optional[Dict[str, Any]] = Field(None, description="Configuration for pod-based deployment")
hybrid_search: bool = Field(False, description="Whether to enable hybrid search")
metric: str = Field("cosine", description="Distance metric for vector similarity")
batch_size: int = Field(100, description="Batch size for operations")
extra_params: Optional[Dict[str, Any]] = Field(None, description="Additional parameters for Pinecone client")
@model_validator(mode="before")
@classmethod
def check_api_key_or_client(cls, values: Dict[str, Any]) -> Dict[str, Any]:
api_key, client = values.get("api_key"), values.get("client")
if not api_key and not client and "PINECONE_API_KEY" not in os.environ:
raise ValueError(
"Either 'api_key' or 'client' must be provided, or PINECONE_API_KEY environment variable must be set."
)
return values
@model_validator(mode="before")
@classmethod
def check_pod_or_serverless(cls, values: Dict[str, Any]) -> Dict[str, Any]:
pod_config, serverless_config = values.get("pod_config"), values.get("serverless_config")
if pod_config and serverless_config:
raise ValueError(
"Both 'pod_config' and 'serverless_config' cannot be specified. Choose one deployment option."
)
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,
}
@@ -14,9 +14,7 @@ class GoogleMatchingEngineConfig(BaseModel):
credentials_path: Optional[str] = Field(None, description="Path to service account credentials file")
vector_search_api_endpoint: Optional[str] = Field(None, description="Vector search API endpoint")
model_config = {
"extra": "forbid"
}
model_config = {"extra": "forbid"}
def __init__(self, **kwargs):
super().__init__(**kwargs)
@@ -26,4 +24,4 @@ class GoogleMatchingEngineConfig(BaseModel):
def model_post_init(self, _context) -> None:
"""Set collection_name to index_id if not provided"""
if self.collection_name is None:
self.collection_name = self.index_id
self.collection_name = self.index_id
+15 -18
View File
@@ -4,26 +4,14 @@ 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:
@@ -35,17 +23,23 @@ class AnthropicLLM(LLMBase):
def generate_response(
self,
messages: List[Dict[str, str]],
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using Anthropic's Claude model based on the provided messages.
Generate a response based on the given messages using Anthropic.
Args:
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
# Extract system message separately
# Separate system message from other messages
system_message = ""
filtered_messages = []
for message in messages:
@@ -62,6 +56,9 @@ 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
+126 -50
View File
@@ -4,26 +4,14 @@ 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:
@@ -37,29 +25,49 @@ class AWSBedrockLLM(LLMBase):
def _format_messages(self, messages: List[Dict[str, str]]) -> str:
"""
Formats a list of messages into a structured prompt for the model.
Formats a list of messages into the required prompt structure for the model.
Args:
messages (List[Dict[str, str]]): A list of dictionaries containing 'role' and 'content'.
messages (List[Dict[str, str]]): A list of dictionaries where each dictionary represents a message.
Each dictionary contains 'role' and 'content' keys.
Returns:
str: A formatted string combining all messages, structured with roles capitalized and separated by newlines.
"""
formatted_messages = [
f"\n\n{msg['role'].capitalize()}: {msg['content']}" for msg in messages
]
formatted_messages = []
for message in messages:
role = message["role"].capitalize()
content = message["content"]
formatted_messages.append(f"\n\n{role}: {content}")
return "".join(formatted_messages) + "\n\nAssistant:"
def _parse_response(self, response) -> str:
def _parse_response(self, response, tools) -> str:
"""
Extracts the generated response from the API response.
Process the response based on whether tools are used or not.
Args:
response: The raw response from the AWS Bedrock API.
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
str: The generated response text.
str or dict: The processed response.
"""
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", "")
@@ -68,21 +76,22 @@ class AWSBedrockLLM(LLMBase):
provider: str,
model: str,
prompt: str,
model_kwargs: Optional[Dict[str, Any]] = None,
model_kwargs: Optional[Dict[str, Any]] = {},
) -> Dict[str, Any]:
"""
Prepares the input dictionary for the specified provider's model.
Prepares the input dictionary for the specified provider's model by mapping and renaming
keys in the input based on the provider's requirements.
Args:
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.
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.
Returns:
Dict[str, Any]: The prepared input dictionary.
Dict[str, Any]: The prepared input dictionary with the correct keys and values for the specified provider.
"""
model_kwargs = model_kwargs or {}
input_body = {"prompt": prompt, **model_kwargs}
provider_mappings = {
@@ -110,35 +119,102 @@ 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 generate_response(self, messages: List[Dict[str, str]]) -> str:
def _convert_tool_format(self, original_tools):
"""
Generates a response using AWS Bedrock based on the provided messages.
Converts a list of tools from their original format to a new standardized format.
Args:
messages (List[Dict[str, str]]): List of message dictionaries containing 'role' and 'content'.
original_tools (list): A list of dictionaries representing the original tools, each containing a 'type' key and corresponding details.
Returns:
str: The generated response text.
list: A list of dictionaries representing the tools in the new standardized format.
"""
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)
new_tools = []
response = self.client.invoke_model(
body=body,
modelId=self.config.model,
accept="application/json",
contentType="application/json",
)
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", []),
}
},
}
}
return self._parse_response(response)
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)
+50 -31
View File
@@ -9,35 +9,17 @@ 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)
# Ensure model name is set; it should match the Azure OpenAI deployment name.
# Model name should match the custom deployment name chosen for it.
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(
@@ -49,20 +31,54 @@ 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using Azure OpenAI based on the provided messages.
Generate a response based on the given messages using Azure OpenAI.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -71,8 +87,11 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+16 -28
View File
@@ -9,38 +9,20 @@ 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)
# Ensure model name is set; it should match the Azure OpenAI deployment name.
# Model name should match the custom deployment name chosen for it.
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,
@@ -56,14 +38,14 @@ class AzureOpenAIStructuredLLM(LLMBase):
response_format: Optional[str] = None,
) -> str:
"""
Generates a response using Azure OpenAI based on the provided messages.
Generate a response based on the given messages using Azure OpenAI.
Args:
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 from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -72,9 +54,15 @@ 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
if tools:
params["tools"] = tools
params["tool_choice"] = tool_choice
response = self.client.chat.completions.create(**params)
return response.choices[0].message.content
return self._parse_response(response, tools)
+2 -6
View File
@@ -4,12 +4,8 @@ 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):
+47 -20
View File
@@ -9,42 +9,64 @@ 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]],
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using DeepSeek based on the provided messages.
Generate a response based on the given messages using DeepSeek.
Args:
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -53,5 +75,10 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+103 -29
View File
@@ -3,7 +3,8 @@ from typing import Dict, List, Optional
try:
import google.generativeai as genai
from google.generativeai import GenerativeModel
from google.generativeai import GenerativeModel, protos
from google.generativeai.types import content_types
except ImportError:
raise ImportError(
"The 'google-generativeai' library is required. Please install it using 'pip install google-generativeai'."
@@ -14,17 +15,7 @@ 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:
@@ -34,25 +25,51 @@ class GeminiLLM(LLMBase):
genai.configure(api_key=api_key)
self.client = GenerativeModel(model_name=self.config.model)
def _reformat_messages(
self, messages: List[Dict[str, str]]
) -> List[Dict[str, str]]:
def _parse_response(self, response, tools):
"""
Reformats messages to match the Gemini API's expected structure.
Process the response based on whether tools are used or not.
Args:
messages (List[Dict[str, str]]): A list of messages with 'role' and 'content' keys.
response: The raw response from API.
tools: The list of tools provided in the request.
Returns:
List[Dict[str, str]]: Reformatted messages in the required format.
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.
"""
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"]
@@ -65,33 +82,90 @@ class GeminiLLM(LLMBase):
return new_messages
def generate_response(
self, messages: List[Dict[str, str]], response_format: Optional[Dict] = None
) -> str:
def _reformat_tools(self, tools: Optional[List[Dict]]):
"""
Generates a response from Gemini based on the given conversation history.
Reformat tools for Gemini.
Args:
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).
tools: The list of tools provided in the request.
Returns:
str: The generated response as text.
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",
):
"""
Generate a response based on the given messages using Gemini.
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".
Returns:
str: The generated response.
"""
params = {
"temperature": self.config.temperature,
"max_output_tokens": self.config.max_tokens,
"top_p": self.config.top_p,
}
if response_format and response_format.get("type") == "json_object":
if response_format is not None and response_format["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 response.candidates[0].content.parts[0].text
return self._parse_response(response, tools)
+46 -20
View File
@@ -5,26 +5,14 @@ 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:
@@ -33,20 +21,54 @@ 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using Groq based on the provided messages.
Generate a response based on the given messages using Groq.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -57,5 +79,9 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+46 -23
View File
@@ -4,50 +4,70 @@ 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using LiteLLM based on the provided messages.
Generate a response based on the given messages using Litellm.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
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,
@@ -58,6 +78,9 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+46 -22
View File
@@ -3,56 +3,77 @@ 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):
"""
Ensures the specified model exists locally. If not, pulls it from Ollama.
Ensure the specified model exists locally. If not, pull 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using Ollama based on the provided messages.
Generate a response based on the given messages using OpenAI.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -66,5 +87,8 @@ class OllamaLLM(LLMBase):
if response_format:
params["format"] = "json"
if tools:
params["tools"] = tools
response = self.client.chat(**params)
return response["message"]["content"]
return self._parse_response(response, tools)
+45 -22
View File
@@ -9,17 +9,7 @@ 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:
@@ -34,27 +24,57 @@ 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response based on the provided messages using OpenAI or OpenRouter.
Generate a response based on the given messages using OpenAI.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -82,6 +102,9 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+6 -17
View File
@@ -9,28 +9,14 @@ 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 generate_response(
@@ -39,7 +25,7 @@ class OpenAIStructuredLLM(LLMBase):
response_format: Optional[str] = None,
) -> str:
"""
Generates a response using OpenAI based on the provided messages.
Generate a response based on the given messages using OpenAI.
Args:
messages (List[Dict[str, str]]): A list of dictionaries, each containing a 'role' and 'content' key.
@@ -47,7 +33,7 @@ class OpenAIStructuredLLM(LLMBase):
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -57,6 +43,9 @@ 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 response.choices[0].message.content
+45 -20
View File
@@ -5,26 +5,14 @@ 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:
@@ -33,20 +21,54 @@ 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: Optional[str] = None,
) -> str:
response_format=None,
tools: Optional[List[Dict]] = None,
tool_choice: str = "auto",
):
"""
Generates a response using TogetherAI based on the provided messages.
Generate a response based on the given messages using TogetherAI.
Args:
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.
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".
Returns:
str: The generated response from the model.
str: The generated response.
"""
params = {
"model": self.config.model,
@@ -57,6 +79,9 @@ 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 response.choices[0].message.content
return self._parse_response(response, tools)
+1 -5
View File
@@ -15,11 +15,7 @@ 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):
+13 -8
View File
@@ -34,7 +34,8 @@ class Memory(MemoryBase):
def __init__(self, config: MemoryConfig = MemoryConfig()):
self.config = config
self.custom_prompt = self.config.custom_prompt
self.custom_fact_extraction_prompt = self.config.custom_fact_extraction_prompt
self.custom_update_memory_prompt = self.config.custom_update_memory_prompt
self.embedding_model = EmbedderFactory.create(self.config.embedder.provider, self.config.embedder.config)
self.vector_store = VectorStoreFactory.create(
self.config.vector_store.provider, self.config.vector_store.config
@@ -70,13 +71,14 @@ class Memory(MemoryBase):
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"]
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,
@@ -176,8 +178,8 @@ class Memory(MemoryBase):
parsed_messages = parse_messages(messages)
if self.custom_prompt:
system_prompt = self.custom_prompt
if self.custom_fact_extraction_prompt:
system_prompt = self.custom_fact_extraction_prompt
user_prompt = f"Input:\n{parsed_messages}"
else:
system_prompt, user_prompt = get_fact_retrieval_messages(parsed_messages)
@@ -203,7 +205,8 @@ class Memory(MemoryBase):
messages_embeddings = self.embedding_model.embed(new_mem, "add")
new_message_embeddings[new_mem] = messages_embeddings
existing_memories = self.vector_store.search(
query=messages_embeddings,
query=new_mem,
vectors=messages_embeddings,
limit=5,
filters=filters,
)
@@ -221,7 +224,9 @@ class Memory(MemoryBase):
temp_uuid_mapping[str(idx)] = item["id"]
retrieved_old_memory[idx]["id"] = str(idx)
function_calling_prompt = get_update_memory_messages(retrieved_old_memory, new_retrieved_facts)
function_calling_prompt = get_update_memory_messages(
retrieved_old_memory, new_retrieved_facts, self.custom_update_memory_prompt
)
try:
new_memories_with_actions = self.llm.generate_response(
@@ -478,7 +483,7 @@ class Memory(MemoryBase):
def _search_vector_store(self, query, filters, limit):
embeddings = self.embedding_model.embed(query, "search")
memories = self.vector_store.search(query=embeddings, limit=limit, filters=filters)
memories = self.vector_store.search(query=query, vectors=embeddings, limit=limit, filters=filters)
excluded_keys = {
"user_id",
+1
View File
@@ -67,6 +67,7 @@ class VectorStoreFactory:
"pgvector": "mem0.vector_stores.pgvector.PGVector",
"milvus": "mem0.vector_stores.milvus.MilvusDB",
"azure_ai_search": "mem0.vector_stores.azure_ai_search.AzureAISearch",
"pinecone": "mem0.vector_stores.pinecone.PineconeDB",
"redis": "mem0.vector_stores.redis.RedisDB",
"elasticsearch": "mem0.vector_stores.elasticsearch.ElasticsearchDB",
"vertex_ai_vector_search": "mem0.vector_stores.vertex_ai_vector_search.GoogleMatchingEngine",
+40 -42
View File
@@ -45,8 +45,10 @@ class AzureAISearch(VectorStoreBase):
collection_name,
api_key,
embedding_model_dims,
compression_type: Optional[str] = None,
compression_type: Optional[str] = None,
use_float16: bool = False,
hybrid_search: bool = False,
vector_filter_mode: Optional[str] = None,
):
"""
Initialize the Azure AI Search vector store.
@@ -60,13 +62,17 @@ class AzureAISearch(VectorStoreBase):
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.)
hybrid_search (bool): Whether to use hybrid search. Default is False.
vector_filter_mode (Optional[str]): Mode for vector filtering. Default is "preFilter".
"""
self.index_name = collection_name
self.collection_name = collection_name
self.embedding_model_dims = embedding_model_dims
# If compression_type is None, treat it as "none".
self.compression_type = (compression_type or "none").lower()
self.compression_type = (compression_type or "none").lower()
self.use_float16 = use_float16
self.hybrid_search = hybrid_search
self.vector_filter_mode = vector_filter_mode
self.search_client = SearchClient(
endpoint=f"https://{service_name}.search.windows.net",
@@ -113,8 +119,6 @@ class AzureAISearch(VectorStoreBase):
)
]
# 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),
@@ -123,11 +127,11 @@ class AzureAISearch(VectorStoreBase):
SearchField(
name="vector",
type=vector_type,
searchable=True,
searchable=True,
vector_search_dimensions=self.embedding_model_dims,
vector_search_profile_name="my-vector-config",
),
SimpleField(name="payload", type=SearchFieldDataType.String, searchable=True),
SearchField(name="payload", type=SearchFieldDataType.String, searchable=True),
]
vector_search = VectorSearch(
@@ -135,7 +139,7 @@ class AzureAISearch(VectorStoreBase):
VectorSearchProfile(
name="my-vector-config",
algorithm_configuration_name="my-algorithms-config",
compression_name=compression_name if self.compression_type != "none" else None
compression_name=compression_name if self.compression_type != "none" else None,
)
],
algorithms=[HnswAlgorithmConfiguration(name="my-algorithms-config")],
@@ -164,12 +168,11 @@ class AzureAISearch(VectorStoreBase):
"""
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)
]
response = self.search_client.upload_documents(documents)
for doc in response:
if not doc.get("status", False):
if not hasattr(doc, "status_code") and doc.get("status_code") != 201:
raise Exception(f"Insert failed for document {doc.get('id')}: {doc}")
return response
@@ -189,16 +192,15 @@ class AzureAISearch(VectorStoreBase):
filter_expression = " and ".join(filter_conditions)
return filter_expression
def search(self, query, limit=5, filters=None, vector_filter_mode="preFilter"):
def search(self, query, vectors, limit=5, filters=None):
"""
Search for similar vectors.
Args:
query (List[float]): Query vector.
query (str): Query.
vectors (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[OutputData]: Search results.
@@ -207,24 +209,28 @@ class AzureAISearch(VectorStoreBase):
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_filter_mode=vector_filter_mode,
)
vector_query = VectorizedQuery(vector=vectors, k_nearest_neighbors=limit, fields="vector")
if self.hybrid_search:
search_results = self.search_client.search(
search_text=query,
vector_queries=[vector_query],
filter=filter_expression,
top=limit,
vector_filter_mode=self.vector_filter_mode,
search_fields=["payload"],
)
else:
search_results = self.search_client.search(
vector_queries=[vector_query],
filter=filter_expression,
top=limit,
vector_filter_mode=self.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):
@@ -236,7 +242,7 @@ class AzureAISearch(VectorStoreBase):
"""
response = self.search_client.delete_documents(documents=[{"id": vector_id}])
for doc in response:
if not doc.get("status", False):
if not hasattr(doc, "status_code") and doc.get("status_code") != 200:
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
@@ -260,7 +266,7 @@ class AzureAISearch(VectorStoreBase):
document[field] = payload.get(field)
response = self.search_client.merge_or_upload_documents(documents=[document])
for doc in response:
if not doc.get("status", False):
if not hasattr(doc, "status_code") and doc.get("status_code") != 200:
raise Exception(f"Update failed for document {vector_id}: {doc}")
return response
@@ -278,9 +284,7 @@ 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]:
"""
@@ -324,18 +328,12 @@ class AzureAISearch(VectorStoreBase):
if filters:
filter_expression = self._build_filter_expression(filters)
search_results = self.search_client.search(
search_text="*", filter=filter_expression, top=limit
)
search_results = self.search_client.search(search_text="*", filter=filter_expression, top=limit)
results = []
for result in search_results:
payload = json.loads(result["payload"])
results.append(
OutputData(
id=result["id"], score=result["@search.score"], payload=payload
)
)
return results
results.append(OutputData(id=result["id"], score=result["@search.score"], payload=payload))
return [results]
def __del__(self):
"""Close the search client when the object is deleted."""
+1 -1
View File
@@ -13,7 +13,7 @@ class VectorStoreBase(ABC):
pass
@abstractmethod
def search(self, query, limit=5, filters=None):
def search(self, query, vectors, limit=5, filters=None):
"""Search for similar vectors."""
pass
+6 -3
View File
@@ -127,19 +127,22 @@ class ChromaDB(VectorStoreBase):
logger.info(f"Inserting {len(vectors)} vectors into collection {self.collection_name}")
self.collection.add(ids=ids, embeddings=vectors, metadatas=payloads)
def search(self, query: List[list], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
def search(
self, query: str, vectors: List[list], limit: int = 5, filters: Optional[Dict] = None
) -> List[OutputData]:
"""
Search for similar vectors.
Args:
query (List[list]): Query vector.
query (str): Query.
vectors (List[list]): List of vectors to search.
limit (int, optional): Number of results to return. Defaults to 5.
filters (Optional[Dict], optional): Filters to apply to the search. Defaults to None.
Returns:
List[OutputData]: Search results.
"""
results = self.collection.query(query_embeddings=query, where=filters, n_results=limit)
results = self.collection.query(query_embeddings=vectors, where=filters, n_results=limit)
final_results = self._parse_output(results)
return final_results
+1
View File
@@ -14,6 +14,7 @@ class VectorStoreConfig(BaseModel):
"qdrant": "QdrantConfig",
"chroma": "ChromaDbConfig",
"pgvector": "PGVectorConfig",
"pinecone": "PineconeConfig",
"milvus": "MilvusDBConfig",
"azure_ai_search": "AzureAISearchConfig",
"redis": "RedisDBConfig",
+21 -17
View File
@@ -46,6 +46,11 @@ class ElasticsearchDB(VectorStoreBase):
if config.auto_create_index:
self.create_index()
if config.custom_search_query:
self.custom_search_query = config.custom_search_query
else:
self.custom_search_query = None
def create_index(self) -> None:
"""Create Elasticsearch index with proper mappings if it doesn't exist"""
index_settings = {
@@ -116,26 +121,25 @@ class ElasticsearchDB(VectorStoreBase):
)
return results
def search(self, query: List[float], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
"""Search for similar vectors using KNN search with pre-filtering."""
if not filters:
# If no filters, just do KNN search
search_query = {"knn": {"field": "vector", "query_vector": query, "k": limit, "num_candidates": limit * 2}}
def search(
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None
) -> List[OutputData]:
"""
Search with two options:
1. Use custom search query if provided
2. Use KNN search on vectors with pre-filtering if no custom search query is provided
"""
if self.custom_search_query:
search_query = self.custom_search_query(vectors, limit, filters)
else:
# If filters exist, apply them with KNN search
filter_conditions = []
for key, value in filters.items():
filter_conditions.append({"term": {f"metadata.{key}": value}})
search_query = {
"knn": {
"field": "vector",
"query_vector": query,
"k": limit,
"num_candidates": limit * 2,
"filter": {"bool": {"must": filter_conditions}},
}
"knn": {"field": "vector", "query_vector": vectors, "k": limit, "num_candidates": limit * 2}
}
if filters:
filter_conditions = []
for key, value in filters.items():
filter_conditions.append({"term": {f"metadata.{key}": value}})
search_query["knn"]["filter"] = {"bool": {"must": filter_conditions}}
response = self.client.search(index=self.collection_name, body=search_query)
+4 -3
View File
@@ -134,12 +134,13 @@ class MilvusDB(VectorStoreBase):
return memory
def search(self, query: list, limit: int = 5, filters: dict = None) -> list:
def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> list:
"""
Search for similar vectors.
Args:
query (List[float]): Query vector.
query (str): Query.
vectors (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.
@@ -149,7 +150,7 @@ class MilvusDB(VectorStoreBase):
query_filter = self._create_filter(filters) if filters else None
hits = self.client.search(
collection_name=self.collection_name,
data=[query],
data=[vectors],
limit=limit,
filter=query_filter,
output_fields=["*"],
+14 -6
View File
@@ -2,7 +2,7 @@ import logging
from typing import Any, Dict, List, Optional
try:
from opensearchpy import OpenSearch
from opensearchpy import OpenSearch, RequestsHttpConnection
from opensearchpy.helpers import bulk
except ImportError:
raise ImportError("OpenSearch requires extra dependencies. Install with `pip install opensearch-py`") from None
@@ -28,9 +28,12 @@ class OpenSearchDB(VectorStoreBase):
# Initialize OpenSearch client
self.client = OpenSearch(
hosts=[{"host": config.host, "port": config.port or 9200}],
http_auth=(config.user, config.password) if (config.user and config.password) else None,
http_auth=config.http_auth
if config.http_auth
else ((config.user, config.password) if (config.user and config.password) else None),
use_ssl=config.use_ssl,
verify_certs=config.verify_certs,
connection_class=RequestsHttpConnection,
)
self.collection_name = config.collection_name
@@ -43,14 +46,17 @@ class OpenSearchDB(VectorStoreBase):
def create_index(self) -> None:
"""Create OpenSearch index with proper mappings if it doesn't exist."""
index_settings = {
# ToDo change replicas to 1
"settings": {
"index": {"number_of_replicas": 1, "number_of_shards": 5, "refresh_interval": "1s", "knn": True}
},
"mappings": {
"properties": {
"text": {"type": "text"},
"vector": {"type": "knn_vector", "dimension": self.vector_dim},
"vector": {
"type": "knn_vector",
"dimension": self.vector_dim,
"method": {"engine": "lucene", "name": "hnsw", "space_type": "cosinesimil"},
},
"metadata": {"type": "object", "properties": {"user_id": {"type": "keyword"}}},
}
},
@@ -111,14 +117,16 @@ class OpenSearchDB(VectorStoreBase):
results.append(OutputData(id=id_, score=1.0, payload=payloads[i]))
return results
def search(self, query: List[float], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
def search(
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None
) -> List[OutputData]:
"""Search for similar vectors using OpenSearch k-NN search with pre-filtering."""
search_query = {
"size": limit,
"query": {
"knn": {
"vector": {
"vector": query,
"vector": vectors,
"k": limit,
}
}
+4 -3
View File
@@ -120,12 +120,13 @@ class PGVector(VectorStoreBase):
)
self.conn.commit()
def search(self, query, limit=5, filters=None):
def search(self, query, vectors, limit=5, filters=None):
"""
Search for similar vectors.
Args:
query (List[float]): Query vector.
query (str): Query.
vectors (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.
@@ -150,7 +151,7 @@ class PGVector(VectorStoreBase):
ORDER BY distance
LIMIT %s
""",
(query, *filter_params, limit),
(vectors, *filter_params, limit),
)
results = self.cur.fetchall()
+369
View File
@@ -0,0 +1,369 @@
import logging
import os
from typing import Any, Dict, List, Optional, Union
from pydantic import BaseModel
try:
from pinecone import Pinecone, PodSpec, ServerlessSpec
from pinecone.data.dataclasses.vector import Vector
except ImportError:
raise ImportError("Pinecone requires extra dependencies. Install with `pip install pinecone pinecone-text`") from None
from mem0.vector_stores.base import VectorStoreBase
logger = logging.getLogger(__name__)
class OutputData(BaseModel):
id: Optional[str] # memory id
score: Optional[float] # distance
payload: Optional[Dict] # metadata
class PineconeDB(VectorStoreBase):
def __init__(
self,
collection_name: str,
embedding_model_dims: int,
client: Optional["Pinecone"],
api_key: Optional[str],
environment: Optional[str],
serverless_config: Optional[Dict[str, Any]],
pod_config: Optional[Dict[str, Any]],
hybrid_search: bool,
metric: str,
batch_size: int,
extra_params: Optional[Dict[str, Any]]
):
"""
Initialize the Pinecone vector store.
Args:
collection_name (str): Name of the index/collection.
embedding_model_dims (int): Dimensions of the embedding model.
client (Pinecone, optional): Existing Pinecone client instance. Defaults to None.
api_key (str, optional): API key for Pinecone. Defaults to None.
environment (str, optional): Pinecone environment. Defaults to None.
serverless_config (Dict, optional): Configuration for serverless deployment. Defaults to None.
pod_config (Dict, optional): Configuration for pod-based deployment. Defaults to None.
hybrid_search (bool, optional): Whether to enable hybrid search. Defaults to False.
metric (str, optional): Distance metric for vector similarity. Defaults to "cosine".
batch_size (int, optional): Batch size for operations. Defaults to 100.
extra_params (Dict, optional): Additional parameters for Pinecone client. Defaults to None.
"""
if client:
self.client = client
else:
api_key = api_key or os.environ.get("PINECONE_API_KEY")
if not api_key:
raise ValueError(
"Pinecone API key must be provided either as a parameter or as an environment variable"
)
params = extra_params or {}
self.client = Pinecone(api_key=api_key, **params)
self.collection_name = collection_name
self.embedding_model_dims = embedding_model_dims
self.environment = environment
self.serverless_config = serverless_config
self.pod_config = pod_config
self.hybrid_search = hybrid_search
self.metric = metric
self.batch_size = batch_size
self.sparse_encoder = None
if self.hybrid_search:
try:
from pinecone_text.sparse import BM25Encoder
logger.info("Initializing BM25Encoder for sparse vectors...")
self.sparse_encoder = BM25Encoder.default()
except ImportError:
logger.warning("pinecone-text not installed. Hybrid search will be disabled.")
self.hybrid_search = False
self.create_col(embedding_model_dims, metric)
def create_col(self, vector_size: int, metric: str = "cosine"):
"""
Create a new index/collection.
Args:
vector_size (int): Size of the vectors to be stored.
metric (str, optional): Distance metric for vector similarity. Defaults to "cosine".
"""
existing_indexes = self.list_cols().names()
if self.collection_name in existing_indexes:
logging.debug(f"Index {self.collection_name} already exists. Skipping creation.")
self.index = self.client.Index(self.collection_name)
return
if self.serverless_config:
spec = ServerlessSpec(**self.serverless_config)
elif self.pod_config:
spec = PodSpec(**self.pod_config)
else:
spec = ServerlessSpec(cloud="aws", region="us-west-2")
self.client.create_index(
name=self.collection_name,
dimension=vector_size,
metric=metric,
spec=spec,
)
self.index = self.client.Index(self.collection_name)
def insert(
self,
vectors: List[List[float]],
payloads: Optional[List[Dict]] = None,
ids: Optional[List[Union[str, int]]] = None,
):
"""
Insert vectors into an index.
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 index {self.collection_name}")
items = []
for idx, vector in enumerate(vectors):
item_id = str(ids[idx]) if ids is not None else str(idx)
payload = payloads[idx] if payloads else {}
vector_record = {"id": item_id, "values": vector, "metadata": payload}
if self.hybrid_search and self.sparse_encoder and "text" in payload:
sparse_vector = self.sparse_encoder.encode_documents(payload["text"])
vector_record["sparse_values"] = sparse_vector
items.append(vector_record)
if len(items) >= self.batch_size:
self.index.upsert(vectors=items)
items = []
if items:
self.index.upsert(vectors=items)
def _parse_output(self, data: Dict) -> List[OutputData]:
"""
Parse the output data from Pinecone search results.
Args:
data (Dict): Output data from Pinecone query.
Returns:
List[OutputData]: Parsed output data.
"""
if isinstance(data, Vector):
result = OutputData(
id=data.id,
score=0.0,
payload=data.metadata,
)
return result
else:
result = []
for match in data:
entry = OutputData(
id=match.get("id"),
score=match.get("score"),
payload=match.get("metadata"),
)
result.append(entry)
return result
def _create_filter(self, filters: Optional[Dict]) -> Dict:
"""
Create a filter dictionary from the provided filters.
"""
if not filters:
return {}
pinecone_filter = {}
for key, value in filters.items():
if isinstance(value, dict) and "gte" in value and "lte" in value:
pinecone_filter[key] = {"$gte": value["gte"], "$lte": value["lte"]}
else:
pinecone_filter[key] = {"$eq": value}
return pinecone_filter
def search(self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
"""
Search for similar vectors.
Args:
query (str): Query.
vectors (list): List of vectors to search.
limit (int, optional): Number of results to return. Defaults to 5.
filters (dict, optional): Filters to apply to the search. Defaults to None.
Returns:
list: Search results.
"""
filter_dict = self._create_filter(filters) if filters else None
query_params = {
"vector": vectors,
"top_k": limit,
"include_metadata": True,
"include_values": False,
}
if filter_dict:
query_params["filter"] = filter_dict
if self.hybrid_search and self.sparse_encoder and "text" in filters:
query_text = filters.get("text")
if query_text:
sparse_vector = self.sparse_encoder.encode_queries(query_text)
query_params["sparse_vector"] = sparse_vector
response = self.index.query(**query_params)
results = self._parse_output(response.matches)
return results
def delete(self, vector_id: Union[str, int]):
"""
Delete a vector by ID.
Args:
vector_id (Union[str, int]): ID of the vector to delete.
"""
self.index.delete(ids=[str(vector_id)])
def update(self, vector_id: Union[str, int], vector: Optional[List[float]] = None, payload: Optional[Dict] = None):
"""
Update a vector and its payload.
Args:
vector_id (Union[str, int]): ID of the vector to update.
vector (list, optional): Updated vector. Defaults to None.
payload (dict, optional): Updated payload. Defaults to None.
"""
item = {
"id": str(vector_id),
}
if vector is not None:
item["values"] = vector
if payload is not None:
item["metadata"] = payload
if self.hybrid_search and self.sparse_encoder and "text" in payload:
sparse_vector = self.sparse_encoder.encode_documents(payload["text"])
item["sparse_values"] = sparse_vector
self.index.upsert(vectors=[item])
def get(self, vector_id: Union[str, int]) -> OutputData:
"""
Retrieve a vector by ID.
Args:
vector_id (Union[str, int]): ID of the vector to retrieve.
Returns:
dict: Retrieved vector or None if not found.
"""
try:
response = self.index.fetch(ids=[str(vector_id)])
if str(vector_id) in response.vectors:
return self._parse_output(response.vectors[str(vector_id)])
return None
except Exception as e:
logger.error(f"Error retrieving vector {vector_id}: {e}")
return None
def list_cols(self):
"""
List all indexes/collections.
Returns:
list: List of index information.
"""
return self.client.list_indexes()
def delete_col(self):
"""Delete an index/collection."""
try:
self.client.delete_index(self.collection_name)
logger.info(f"Index {self.collection_name} deleted successfully")
except Exception as e:
logger.error(f"Error deleting index {self.collection_name}: {e}")
def col_info(self) -> Dict:
"""
Get information about an index/collection.
Returns:
dict: Index information.
"""
return self.client.describe_index(self.collection_name)
def list(self, filters: Optional[Dict] = None, limit: int = 100) -> List[OutputData]:
"""
List vectors in an index with optional filtering.
Args:
filters (dict, optional): Filters to apply to the list. Defaults to None.
limit (int, optional): Number of vectors to return. Defaults to 100.
Returns:
dict: List of vectors with their metadata.
"""
filter_dict = self._create_filter(filters) if filters else None
stats = self.index.describe_index_stats()
dimension = stats.dimension
zero_vector = [0.0] * dimension
query_params = {
"vector": zero_vector,
"top_k": limit,
"include_metadata": True,
"include_values": True,
}
if filter_dict:
query_params["filter"] = filter_dict
try:
response = self.index.query(**query_params)
response = response.to_dict()
results = self._parse_output(response["matches"])
return [results]
except Exception as e:
logger.error(f"Error listing vectors: {e}")
return {"points": [], "next_page_token": None}
def count(self) -> int:
"""
Count number of vectors in the index.
Returns:
int: Total number of vectors.
"""
stats = self.index.describe_index_stats()
return stats.total_vector_count
def reset(self):
"""
Reset the index by deleting and recreating it.
"""
self.delete_col()
self.create_col(self.embedding_model_dims, self.metric)
+4 -3
View File
@@ -127,12 +127,13 @@ class Qdrant(VectorStoreBase):
conditions.append(FieldCondition(key=key, match=MatchValue(value=value)))
return Filter(must=conditions) if conditions else None
def search(self, query: list, limit: int = 5, filters: dict = None) -> list:
def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None) -> list:
"""
Search for similar vectors.
Args:
query (list): Query vector.
query (str): Query.
vectors (list): 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.
@@ -142,7 +143,7 @@ class Qdrant(VectorStoreBase):
query_filter = self._create_filter(filters) if filters else None
hits = self.client.query_points(
collection_name=self.collection_name,
query=query,
query=vectors,
query_filter=query_filter,
limit=limit,
)
+2 -2
View File
@@ -101,12 +101,12 @@ class RedisDB(VectorStoreBase):
data.append(entry)
self.index.load(data, id_field="memory_id")
def search(self, query: list, limit: int = 5, filters: dict = None):
def search(self, query: str, vectors: list, limit: int = 5, filters: dict = None):
conditions = [Tag(key) == value for key, value in filters.items() if value is not None]
filter = reduce(lambda x, y: x & y, conditions)
v = VectorQuery(
vector=np.array(query, dtype=np.float32).tobytes(),
vector=np.array(vectors, dtype=np.float32).tobytes(),
vector_field_name="embedding",
return_fields=["memory_id", "hash", "agent_id", "run_id", "user_id", "memory", "metadata", "created_at"],
filter_expression=filter,
+6 -6
View File
@@ -112,16 +112,18 @@ class Supabase(VectorStoreBase):
payloads = [{} for _ in vectors]
records = [(id, vector, payload) for id, vector, payload in zip(ids, vectors, payloads)]
print(records)
self.collection.upsert(records)
def search(self, query: List[float], limit: int = 5, filters: Optional[dict] = None) -> List[OutputData]:
def search(
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[dict] = None
) -> List[OutputData]:
"""
Search for similar vectors.
Args:
query (List[float]): Query vector
query (str): Query.
vectors (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.
@@ -129,11 +131,9 @@ class Supabase(VectorStoreBase):
List[OutputData]: Search results
"""
filters = self._preprocess_filters(filters)
print(filters)
results = self.collection.query(
data=query, limit=limit, filters=filters, include_metadata=True, include_value=True
data=vectors, limit=limit, filters=filters, include_metadata=True, include_value=True
)
print(results)
return [OutputData(id=str(result[0]), score=float(result[1]), payload=result[2]) for result in results]
+116 -181
View File
@@ -32,19 +32,19 @@ class GoogleMatchingEngine(VectorStoreBase):
def __init__(self, **kwargs):
"""Initialize Google Matching Engine client."""
logger.debug("Initializing Google Matching Engine with kwargs: %s", kwargs)
# If collection_name is passed, use it as deployment_index_id if deployment_index_id is not provided
if 'collection_name' in kwargs and 'deployment_index_id' not in kwargs:
kwargs['deployment_index_id'] = kwargs['collection_name']
logger.debug("Using collection_name as deployment_index_id: %s", kwargs['deployment_index_id'])
elif 'deployment_index_id' in kwargs and 'collection_name' not in kwargs:
kwargs['collection_name'] = kwargs['deployment_index_id']
logger.debug("Using deployment_index_id as collection_name: %s", kwargs['collection_name'])
if "collection_name" in kwargs and "deployment_index_id" not in kwargs:
kwargs["deployment_index_id"] = kwargs["collection_name"]
logger.debug("Using collection_name as deployment_index_id: %s", kwargs["deployment_index_id"])
elif "deployment_index_id" in kwargs and "collection_name" not in kwargs:
kwargs["collection_name"] = kwargs["deployment_index_id"]
logger.debug("Using deployment_index_id as collection_name: %s", kwargs["collection_name"])
try:
config = GoogleMatchingEngineConfig(**kwargs)
logger.debug("Config created: %s", config.model_dump())
logger.debug("Config collection_name: %s", getattr(config, 'collection_name', None))
logger.debug("Config collection_name: %s", getattr(config, "collection_name", None))
except Exception as e:
logger.error("Failed to validate config: %s", str(e))
raise
@@ -57,41 +57,37 @@ class GoogleMatchingEngine(VectorStoreBase):
self.deployment_index_id = config.deployment_index_id # The deployment-specific ID
self.collection_name = config.collection_name
self.vector_search_api_endpoint = config.vector_search_api_endpoint
logger.debug("Using project=%s, location=%s", self.project_id, self.region)
# Initialize Vertex AI with credentials if provided
init_args = {
"project": self.project_id,
"location": self.region,
}
if hasattr(config, 'credentials_path') and config.credentials_path:
if hasattr(config, "credentials_path") and config.credentials_path:
logger.debug("Using credentials from: %s", config.credentials_path)
credentials = service_account.Credentials.from_service_account_file(
config.credentials_path
)
credentials = service_account.Credentials.from_service_account_file(config.credentials_path)
init_args["credentials"] = credentials
try:
aiplatform.init(**init_args)
logger.debug("Vertex AI initialized successfully")
except Exception as e:
logger.error("Failed to initialize Vertex AI: %s", str(e))
raise
try:
# Format the index path properly using the configured index_id
index_path = f"projects/{self.project_number}/locations/{self.region}/indexes/{self.index_id}"
logger.debug("Initializing index with path: %s", index_path)
self.index = aiplatform.MatchingEngineIndex(index_name=index_path)
logger.debug("Index initialized successfully")
# Format the endpoint name properly
endpoint_name = self.endpoint_id
logger.debug("Initializing endpoint with name: %s", endpoint_name)
self.index_endpoint = aiplatform.MatchingEngineIndexEndpoint(
index_endpoint_name=endpoint_name
)
self.index_endpoint = aiplatform.MatchingEngineIndexEndpoint(index_endpoint_name=endpoint_name)
logger.debug("Endpoint initialized successfully")
except Exception as e:
logger.error("Failed to initialize Matching Engine components: %s", str(e))
@@ -119,47 +115,36 @@ class GoogleMatchingEngine(VectorStoreBase):
def _create_restriction(self, key: str, value: Any) -> aiplatform_v1.types.index.IndexDatapoint.Restriction:
"""Create a restriction object for the Matching Engine index.
Args:
key: The namespace/key for the restriction
value: The value to restrict on
Returns:
Restriction object for the index
"""
str_value = str(value) if value is not None else ""
return aiplatform_v1.types.index.IndexDatapoint.Restriction(
namespace=key,
allow_list=[str_value]
)
return aiplatform_v1.types.index.IndexDatapoint.Restriction(namespace=key, allow_list=[str_value])
def _create_datapoint(
self,
vector_id: str,
vector: List[float],
payload: Optional[Dict] = None
self, vector_id: str, vector: List[float], payload: Optional[Dict] = None
) -> aiplatform_v1.types.index.IndexDatapoint:
"""Create a datapoint object for the Matching Engine index.
Args:
vector_id: The ID for the datapoint
vector: The vector to store
payload: Optional metadata to store with the vector
Returns:
IndexDatapoint object
"""
restrictions = []
if payload:
restrictions = [
self._create_restriction(key, value)
for key, value in payload.items()
]
restrictions = [self._create_restriction(key, value) for key, value in payload.items()]
return aiplatform_v1.types.index.IndexDatapoint(
datapoint_id=vector_id,
feature_vector=vector,
restricts=restrictions
datapoint_id=vector_id, feature_vector=vector, restricts=restrictions
)
def insert(
@@ -169,41 +154,41 @@ class GoogleMatchingEngine(VectorStoreBase):
ids: Optional[List[str]] = None,
) -> None:
"""Insert vectors into the Matching Engine index.
Args:
vectors: List of vectors to insert
payloads: Optional list of metadata dictionaries
ids: Optional list of IDs for the vectors
Raises:
ValueError: If vectors is empty or lengths don't match
GoogleAPIError: If the API call fails
"""
if not vectors:
raise ValueError("No vectors provided for insertion")
if payloads and len(payloads) != len(vectors):
raise ValueError(f"Number of payloads ({len(payloads)}) does not match number of vectors ({len(vectors)})")
if ids and len(ids) != len(vectors):
raise ValueError(f"Number of ids ({len(ids)}) does not match number of vectors ({len(vectors)})")
logger.debug("Starting insert of %d vectors", len(vectors))
try:
datapoints = [
self._create_datapoint(
vector_id=ids[i] if ids else str(uuid.uuid4()),
vector=vector,
payload=payloads[i] if payloads and i < len(payloads) else None
payload=payloads[i] if payloads and i < len(payloads) else None,
)
for i, vector in enumerate(vectors)
]
logger.debug("Created %d datapoints", len(datapoints))
self.index.upsert_datapoints(datapoints=datapoints)
logger.debug("Successfully inserted datapoints")
except google.api_core.exceptions.GoogleAPIError as e:
logger.error("Failed to insert vectors: %s", str(e))
raise
@@ -212,21 +197,22 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.error("Stack trace: %s", traceback.format_exc())
raise
def search(self, query: List[float], limit: int = 5, filters: Optional[Dict] = None) -> List[OutputData]:
def search(
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None
) -> List[OutputData]:
"""
Search for similar vectors.
Args:
query (List[float]): Query vector.
query (str): Query.
vectors (List[float]): Query vector.
limit (int, optional): Number of results to return. Defaults to 5.
filters (Optional[Dict], optional): Filters to apply to the search. Defaults to None.
Returns:
List[OutputData]: Search results (unwrapped)
"""
logger.debug("Starting search")
logger.debug("Query type: %s, length: %d", type(query), len(query))
logger.debug("Limit: %d, Filters: %s", limit, filters)
try:
filter_namespaces = []
if filters:
@@ -235,53 +221,42 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.debug("Processing filter %s=%s (type=%s)", key, value, type(value))
if isinstance(value, (str, int, float)):
logger.debug("Adding simple filter for %s", key)
filter_namespaces.append(
Namespace(key, [str(value)], [])
)
filter_namespaces.append(Namespace(key, [str(value)], []))
elif isinstance(value, dict):
logger.debug("Adding complex filter for %s", key)
includes = value.get('include', [])
excludes = value.get('exclude', [])
filter_namespaces.append(
Namespace(key, includes, excludes)
)
includes = value.get("include", [])
excludes = value.get("exclude", [])
filter_namespaces.append(Namespace(key, includes, excludes))
logger.debug("Final filter_namespaces: %s", filter_namespaces)
response = self.index_endpoint.find_neighbors(
deployed_index_id=self.deployment_index_id,
queries=[query],
queries=[vectors],
num_neighbors=limit,
filter=filter_namespaces if filter_namespaces else None,
return_full_datapoint=True
return_full_datapoint=True,
)
if not response or len(response) == 0 or len(response[0]) == 0:
logger.debug("No results found")
return []
results = []
for neighbor in response[0]:
logger.debug("Processing neighbor - id: %s, distance: %s",
neighbor.id, neighbor.distance)
logger.debug("Processing neighbor - id: %s, distance: %s", neighbor.id, neighbor.distance)
payload = {}
if hasattr(neighbor, 'restricts'):
if hasattr(neighbor, "restricts"):
logger.debug("Processing restricts")
for restrict in neighbor.restricts:
if (hasattr(restrict, 'name') and
hasattr(restrict, 'allow_tokens') and
restrict.allow_tokens):
if hasattr(restrict, "name") and hasattr(restrict, "allow_tokens") and restrict.allow_tokens:
logger.debug("Adding %s: %s", restrict.name, restrict.allow_tokens[0])
payload[restrict.name] = restrict.allow_tokens[0]
output_data = OutputData(
id=neighbor.id,
score=neighbor.distance,
payload=payload
)
output_data = OutputData(id=neighbor.id, score=neighbor.distance, payload=payload)
results.append(output_data)
logger.debug("Returning %d results", len(results))
return results
@@ -291,7 +266,6 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.error("Stack trace: %s", traceback.format_exc())
raise
def delete(self, vector_id: Optional[str] = None, ids: Optional[List[str]] = None) -> bool:
"""
Delete vectors from the Matching Engine index.
@@ -326,14 +300,13 @@ class GoogleMatchingEngine(VectorStoreBase):
except google.api_core.exceptions.InvalidArgument as e:
logger.error("Invalid argument: %s", str(e))
return False
except Exception as e:
logger.error("Error occurred: %s", str(e))
logger.error("Error type: %s", type(e))
logger.error("Stack trace: %s", traceback.format_exc())
return False
def update(
self,
vector_id: str,
@@ -341,42 +314,40 @@ class GoogleMatchingEngine(VectorStoreBase):
payload: Optional[Dict] = None,
) -> bool:
"""Update a vector and its payload.
Args:
vector_id: ID of the vector to update
vector: Optional new vector values
payload: Optional new metadata payload
Returns:
bool: True if update was successful
Raises:
ValueError: If neither vector nor payload is provided
GoogleAPIError: If the API call fails
"""
logger.debug("Starting update for vector_id: %s", vector_id)
if vector is None and payload is None:
raise ValueError("Either vector or payload must be provided for update")
# First check if the vector exists
try:
existing = self.get(vector_id)
if existing is None:
logger.error("Vector ID not found: %s", vector_id)
return False
datapoint = self._create_datapoint(
vector_id=vector_id,
vector=vector if vector is not None else [],
payload=payload
vector_id=vector_id, vector=vector if vector is not None else [], payload=payload
)
logger.debug("Upserting datapoint: %s", datapoint)
self.index.upsert_datapoints(datapoints=[datapoint])
logger.debug("Update completed successfully")
return True
except google.api_core.exceptions.GoogleAPIError as e:
logger.error("API error during update: %s", str(e))
return False
@@ -385,7 +356,6 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.error("Stack trace: %s", traceback.format_exc())
raise
def get(self, vector_id: str) -> Optional[OutputData]:
"""
Retrieve a vector by ID.
@@ -395,24 +365,17 @@ class GoogleMatchingEngine(VectorStoreBase):
Optional[OutputData]: Retrieved vector or None if not found.
"""
logger.debug("Starting get for vector_id: %s", vector_id)
try:
if not self.vector_search_api_endpoint:
raise ValueError("vector_search_api_endpoint is required for get operation")
vector_search_client = aiplatform_v1.MatchServiceClient(
client_options={
"api_endpoint": self.vector_search_api_endpoint
},
)
datapoint = aiplatform_v1.IndexDatapoint(
datapoint_id=vector_id
client_options={"api_endpoint": self.vector_search_api_endpoint},
)
datapoint = aiplatform_v1.IndexDatapoint(datapoint_id=vector_id)
query = aiplatform_v1.FindNeighborsRequest.Query(
datapoint=datapoint,
neighbor_count=1
)
query = aiplatform_v1.FindNeighborsRequest.Query(datapoint=datapoint, neighbor_count=1)
request = aiplatform_v1.FindNeighborsRequest(
index_endpoint=f"projects/{self.project_number}/locations/{self.region}/indexEndpoints/{self.endpoint_id}",
deployed_index_id=self.deployment_index_id,
@@ -423,41 +386,36 @@ class GoogleMatchingEngine(VectorStoreBase):
try:
response = vector_search_client.find_neighbors(request)
logger.debug("Got response")
if response and response.nearest_neighbors:
nearest = response.nearest_neighbors[0]
if nearest.neighbors:
neighbor = nearest.neighbors[0]
payload = {}
if hasattr(neighbor.datapoint, 'restricts'):
if hasattr(neighbor.datapoint, "restricts"):
for restrict in neighbor.datapoint.restricts:
if restrict.allow_list:
payload[restrict.namespace] = restrict.allow_list[0]
return OutputData(
id=neighbor.datapoint.datapoint_id,
score=neighbor.distance,
payload=payload
)
return OutputData(id=neighbor.datapoint.datapoint_id, score=neighbor.distance, payload=payload)
logger.debug("No results found")
return None
except google.api_core.exceptions.NotFound:
logger.debug("Datapoint not found")
return None
except google.api_core.exceptions.PermissionDenied as e:
logger.error("Permission denied: %s", str(e))
return None
except Exception as e:
logger.error("Error occurred: %s", str(e))
logger.error("Error type: %s", type(e))
logger.error("Stack trace: %s", traceback.format_exc())
raise
def list_cols(self) -> List[str]:
"""
List all collections (indexes).
@@ -466,7 +424,6 @@ class GoogleMatchingEngine(VectorStoreBase):
"""
return [self.deployment_index_id]
def delete_col(self):
"""
Delete a collection (index).
@@ -475,7 +432,6 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.warning("Delete collection operation is not supported for Google Matching Engine")
pass
def col_info(self) -> Dict:
"""
Get information about a collection (index).
@@ -486,17 +442,16 @@ class GoogleMatchingEngine(VectorStoreBase):
"index_id": self.index_id,
"endpoint_id": self.endpoint_id,
"project_id": self.project_id,
"region": self.region
"region": self.region,
}
def list(self, filters: Optional[Dict] = None, limit: Optional[int] = None) -> List[List[OutputData]]:
"""List vectors matching the given filters.
Args:
filters: Optional filters to apply
limit: Optional maximum number of results to return
Returns:
List[List[OutputData]]: List of matching vectors wrapped in an extra array
to match the interface
@@ -504,36 +459,31 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.debug("Starting list operation")
logger.debug("Filters: %s", filters)
logger.debug("Limit: %s", limit)
try:
# Use a zero vector for the search
dimension = 768 # This should be configurable based on the model
zero_vector = [0.0] * dimension
# Use a large limit if none specified
search_limit = limit if limit is not None else 10000
results = self.search(
query=zero_vector,
limit=search_limit,
filters=filters
)
results = self.search(query=zero_vector, limit=search_limit, filters=filters)
logger.debug("Found %d results", len(results))
return [results] # Wrap in extra array to match interface
except Exception as e:
logger.error("Error in list operation: %s", str(e))
logger.error("Stack trace: %s", traceback.format_exc())
raise
def create_col(self, name=None, vector_size=None, distance=None):
"""
Create a new collection. For Google Matching Engine, collections (indexes)
Create a new collection. For Google Matching Engine, collections (indexes)
are created through the Google Cloud Console or API separately.
This method is a no-op since indexes are pre-created.
Args:
name: Ignored for Google Matching Engine
vector_size: Ignored for Google Matching Engine
@@ -543,41 +493,35 @@ class GoogleMatchingEngine(VectorStoreBase):
# This method is included only to satisfy the abstract base class
pass
def add(self, text: str, metadata: Optional[Dict] = None, user_id: Optional[str] = None) -> str:
logger.debug("Starting add operation")
logger.debug("Text: %s", text)
logger.debug("Metadata: %s", metadata)
logger.debug("User ID: %s", user_id)
try:
# Generate a unique ID for this entry
vector_id = str(uuid.uuid4())
# Create the payload with all necessary fields
payload = {
"data": text, # Store the text in the data field
"user_id": user_id,
**(metadata or {})
**(metadata or {}),
}
# Get the embedding
vector = self.embedder.embed_query(text)
# Insert using the insert method
self.insert(
vectors=[vector],
payloads=[payload],
ids=[vector_id]
)
self.insert(vectors=[vector], payloads=[payload], ids=[vector_id])
return vector_id
except Exception as e:
logger.error("Error occurred: %s", str(e))
raise
def add_texts(
self,
texts: List[str],
@@ -585,47 +529,45 @@ class GoogleMatchingEngine(VectorStoreBase):
ids: Optional[List[str]] = None,
) -> List[str]:
"""Add texts to the vector store.
Args:
texts: List of texts to add
metadatas: Optional list of metadata dicts
ids: Optional list of IDs to use
Returns:
List[str]: List of IDs of the added texts
Raises:
ValueError: If texts is empty or lengths don't match
"""
if not texts:
raise ValueError("No texts provided")
if metadatas and len(metadatas) != len(texts):
raise ValueError(f"Number of metadata items ({len(metadatas)}) does not match number of texts ({len(texts)})")
raise ValueError(
f"Number of metadata items ({len(metadatas)}) does not match number of texts ({len(texts)})"
)
if ids and len(ids) != len(texts):
raise ValueError(f"Number of ids ({len(ids)}) does not match number of texts ({len(texts)})")
logger.debug("Starting add_texts operation")
logger.debug("Number of texts: %d", len(texts))
logger.debug("Has metadatas: %s", metadatas is not None)
logger.debug("Has ids: %s", ids is not None)
if ids is None:
ids = [str(uuid.uuid4()) for _ in texts]
try:
# Get embeddings
embeddings = self.embedder.embed_documents(texts)
# Add to store
self.insert(
vectors=embeddings,
payloads=metadatas if metadatas else [{}] * len(texts),
ids=ids
)
self.insert(vectors=embeddings, payloads=metadatas if metadatas else [{}] * len(texts), ids=ids)
return ids
except Exception as e:
logger.error("Error in add_texts: %s", str(e))
logger.error("Stack trace: %s", traceback.format_exc())
@@ -657,18 +599,12 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.debug("Query: %s", query)
logger.debug("k: %d", k)
logger.debug("Filter: %s", filter)
embedding = self.embedder.embed_query(query)
results = self.search(query=embedding, limit=k, filters=filter)
docs_and_scores = [
(
Document(
page_content=result.payload.get("text", ""),
metadata=result.payload
),
result.score
)
(Document(page_content=result.payload.get("text", ""), metadata=result.payload), result.score)
for result in results
]
logger.debug("Found %d results", len(docs_and_scores))
@@ -684,4 +620,3 @@ class GoogleMatchingEngine(VectorStoreBase):
logger.debug("Starting similarity search")
docs_and_scores = self.similarity_search_with_score(query, k, filter)
return [doc for doc, _ in docs_and_scores]
+4 -2
View File
@@ -154,7 +154,9 @@ class Weaviate(VectorStoreBase):
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]:
def search(
self, query: str, vectors: List[float], limit: int = 5, filters: Optional[Dict] = None
) -> List[OutputData]:
"""
Search for similar vectors.
"""
@@ -167,7 +169,7 @@ class Weaviate(VectorStoreBase):
combined_filter = Filter.all_of(filter_conditions) if filter_conditions else None
response = collection.query.hybrid(
query="",
vector=query,
vector=vectors,
limit=limit,
filters=combined_filter,
return_properties=["hash", "created_at", "updated_at", "user_id", "agent_id", "run_id", "data", "category"],
Generated
+251 -32
View File
@@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 1.8.4 and should not be changed by hand.
# This file is automatically @generated by Poetry 2.1.1 and should not be changed by hand.
[[package]]
name = "aiohappyeyeballs"
@@ -6,6 +6,8 @@ version = "2.4.6"
description = "Happy Eyeballs for asyncio"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "aiohappyeyeballs-2.4.6-py3-none-any.whl", hash = "sha256:147ec992cf873d74f5062644332c539fcd42956dc69453fe5204195e560517e1"},
{file = "aiohappyeyeballs-2.4.6.tar.gz", hash = "sha256:9b05052f9042985d32ecbe4b59a77ae19c006a78f1344d7fdad69d28ded3d0b0"},
@@ -17,6 +19,8 @@ version = "3.11.13"
description = "Async http client/server framework (asyncio)"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "aiohttp-3.11.13-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:a4fe27dbbeec445e6e1291e61d61eb212ee9fed6e47998b27de71d70d3e8777d"},
{file = "aiohttp-3.11.13-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:9e64ca2dbea28807f8484c13f684a2f761e69ba2640ec49dacd342763cc265ef"},
@@ -120,6 +124,8 @@ version = "1.3.2"
description = "aiosignal: a list of registered asynchronous callbacks"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "aiosignal-1.3.2-py2.py3-none-any.whl", hash = "sha256:45cde58e409a301715980c2b01d0c28bdde3770d8290b5eb2173759d9acb31a5"},
{file = "aiosignal-1.3.2.tar.gz", hash = "sha256:a8c255c66fafb1e499c9351d0bf32ff2d8a0321595ebac3b93713656d2436f54"},
@@ -134,6 +140,7 @@ version = "0.7.0"
description = "Reusable constraint types to use with typing.Annotated"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "annotated_types-0.7.0-py3-none-any.whl", hash = "sha256:1f02e8b43a8fbbc3f3e0d4f0f4bfc8131bcb4eebe8849b8e5c773f3a1c582a53"},
{file = "annotated_types-0.7.0.tar.gz", hash = "sha256:aff07c09a53a08bc8cfccb9c85b05f1aa9a2a6f23728d790723543408344ce89"},
@@ -145,6 +152,7 @@ version = "4.8.0"
description = "High level compatibility layer for multiple asynchronous event loop implementations"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "anyio-4.8.0-py3-none-any.whl", hash = "sha256:b5011f270ab5eb0abf13385f851315585cc37ef330dd88e27ec3d34d651fd47a"},
{file = "anyio-4.8.0.tar.gz", hash = "sha256:1d9fe889df5212298c0c0723fa20479d1b94883a2df44bd3897aa91083316f7a"},
@@ -167,6 +175,8 @@ version = "4.0.3"
description = "Timeout context manager for asyncio programs"
optional = false
python-versions = ">=3.7"
groups = ["main"]
markers = "extra == \"graph\" and python_version < \"3.11\""
files = [
{file = "async-timeout-4.0.3.tar.gz", hash = "sha256:4640d96be84d82d02ed59ea2b7105a0f7b33abe8703703cd0ab0bf87c427522f"},
{file = "async_timeout-4.0.3-py3-none-any.whl", hash = "sha256:7405140ff1230c310e51dc27b3145b9092d659ce68ff733fb0cefe3ee42be028"},
@@ -178,6 +188,8 @@ version = "25.1.0"
description = "Classes Without Boilerplate"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "attrs-25.1.0-py3-none-any.whl", hash = "sha256:c75a69e28a550a7e93789579c22aa26b0f5b83b75dc4e08fe092980051e1090a"},
{file = "attrs-25.1.0.tar.gz", hash = "sha256:1c97078a80c814273a76b2a298a932eb681c87415c11dee0a6921de7f1b02c3e"},
@@ -247,6 +259,7 @@ version = "2.2.1"
description = "Function decoration for backoff and retry"
optional = false
python-versions = ">=3.7,<4.0"
groups = ["main"]
files = [
{file = "backoff-2.2.1-py3-none-any.whl", hash = "sha256:63579f9a0628e06278f7e47b7d7d5b6ce20dc65c5e96a6f3ca99a6adca0396e8"},
{file = "backoff-2.2.1.tar.gz", hash = "sha256:03f829f5bb1923180821643f8753b0502c3b682293992485b0eef2807afa5cba"},
@@ -258,6 +271,7 @@ version = "2025.1.31"
description = "Python package for providing Mozilla's CA Bundle."
optional = false
python-versions = ">=3.6"
groups = ["main"]
files = [
{file = "certifi-2025.1.31-py3-none-any.whl", hash = "sha256:ca78db4565a652026a4db2bcdf68f2fb589ea80d0be70e03929ed730746b84fe"},
{file = "certifi-2025.1.31.tar.gz", hash = "sha256:3d5da6925056f6f18f119200434a4780a94263f10d1c21d032a6f6b2baa20651"},
@@ -269,6 +283,8 @@ version = "1.17.1"
description = "Foreign Function Interface for Python calling C code."
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\" and platform_python_implementation == \"PyPy\""
files = [
{file = "cffi-1.17.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:df8b1c11f177bc2313ec4b2d46baec87a5f3e71fc8b45dab2ee7cae86d9aba14"},
{file = "cffi-1.17.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:8f2cdc858323644ab277e9bb925ad72ae0e67f69e804f4898c070998d50b1a67"},
@@ -348,6 +364,7 @@ version = "3.4.1"
description = "The Real First Universal Charset Detector. Open, modern and actively maintained alternative to Chardet."
optional = false
python-versions = ">=3.7"
groups = ["main"]
files = [
{file = "charset_normalizer-3.4.1-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:91b36a978b5ae0ee86c394f5a54d6ef44db1de0815eb43de826d41d21e4af3de"},
{file = "charset_normalizer-3.4.1-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7461baadb4dc00fd9e0acbe254e3d7d2112e7f92ced2adc96e54ef6501c5f176"},
@@ -449,10 +466,12 @@ version = "0.4.6"
description = "Cross-platform colored terminal text."
optional = false
python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,>=2.7"
groups = ["main", "dev", "test"]
files = [
{file = "colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6"},
{file = "colorama-0.4.6.tar.gz", hash = "sha256:08695f5cb7ed6e0531a20572697297273c47b8cae5a63ffc6d6ed5c201be6e44"},
]
markers = {main = "platform_system == \"Windows\"", dev = "sys_platform == \"win32\"", test = "sys_platform == \"win32\""}
[[package]]
name = "distro"
@@ -460,6 +479,7 @@ version = "1.9.0"
description = "Distro - an OS platform information API"
optional = false
python-versions = ">=3.6"
groups = ["main"]
files = [
{file = "distro-1.9.0-py3-none-any.whl", hash = "sha256:7bffd925d65168f85027d8da9af6bddab658135b840670a223589bc0c8ef02b2"},
{file = "distro-1.9.0.tar.gz", hash = "sha256:2fa77c6fd8940f116ee1d6b94a2f90b13b5ea8d019b98bc8bafdcabcdd9bdbed"},
@@ -471,6 +491,8 @@ version = "1.2.2"
description = "Backport of PEP 654 (exception groups)"
optional = false
python-versions = ">=3.7"
groups = ["main", "dev", "test"]
markers = "python_version < \"3.11\""
files = [
{file = "exceptiongroup-1.2.2-py3-none-any.whl", hash = "sha256:3111b9d131c238bec2f8f516e123e14ba243563fb135d3fe885990585aa7795b"},
{file = "exceptiongroup-1.2.2.tar.gz", hash = "sha256:47c2edf7c6738fafb49fd34290706d1a1a2f4d1c6df275526b62cbb4aa5393cc"},
@@ -485,6 +507,8 @@ version = "1.5.0"
description = "A list-like structure which implements collections.abc.MutableSequence"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "frozenlist-1.5.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:5b6a66c18b5b9dd261ca98dffcb826a525334b2f29e7caa54e182255c5f6a65a"},
{file = "frozenlist-1.5.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:d1b3eb7b05ea246510b43a7e53ed1653e55c2121019a97e60cad7efb881a97bb"},
@@ -586,6 +610,8 @@ version = "2024.12.0"
description = "File-system specification"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "fsspec-2024.12.0-py3-none-any.whl", hash = "sha256:b520aed47ad9804237ff878b504267a3b0b441e97508bd6d2d8774e3db85cee2"},
{file = "fsspec-2024.12.0.tar.gz", hash = "sha256:670700c977ed2fb51e0d9f9253177ed20cbde4a3e5c0283cc5385b5870c8533f"},
@@ -625,6 +651,8 @@ version = "3.1.1"
description = "Lightweight in-process concurrent programming"
optional = false
python-versions = ">=3.7"
groups = ["main"]
markers = "(platform_machine == \"aarch64\" or platform_machine == \"ppc64le\" or platform_machine == \"x86_64\" or platform_machine == \"amd64\" or platform_machine == \"AMD64\" or platform_machine == \"win32\" or platform_machine == \"WIN32\") and python_version < \"3.14\""
files = [
{file = "greenlet-3.1.1-cp310-cp310-macosx_11_0_universal2.whl", hash = "sha256:0bbae94a29c9e5c7e4a2b7f0aae5c17e8e90acbfd3bf6270eeba60c39fce3563"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:0fde093fb93f35ca72a556cf72c92ea3ebfda3d79fc35bb19fbe685853869a83"},
@@ -711,6 +739,7 @@ version = "1.70.0"
description = "HTTP/2-based RPC framework"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "grpcio-1.70.0-cp310-cp310-linux_armv7l.whl", hash = "sha256:95469d1977429f45fe7df441f586521361e235982a0b39e33841549143ae2851"},
{file = "grpcio-1.70.0-cp310-cp310-macosx_12_0_universal2.whl", hash = "sha256:ed9718f17fbdb472e33b869c77a16d0b55e166b100ec57b016dc7de9c8d236bf"},
@@ -778,6 +807,7 @@ version = "1.70.0"
description = "Protobuf code generator for gRPC"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "grpcio_tools-1.70.0-cp310-cp310-linux_armv7l.whl", hash = "sha256:4d456521290e25b1091975af71604facc5c7db162abdca67e12a0207b8bbacbe"},
{file = "grpcio_tools-1.70.0-cp310-cp310-macosx_12_0_universal2.whl", hash = "sha256:d50080bca84f53f3a05452e06e6251cbb4887f5a1d1321d1989e26d6e0dc398d"},
@@ -847,6 +877,7 @@ version = "0.14.0"
description = "A pure-Python, bring-your-own-I/O implementation of HTTP/1.1"
optional = false
python-versions = ">=3.7"
groups = ["main"]
files = [
{file = "h11-0.14.0-py3-none-any.whl", hash = "sha256:e3fe4ac4b851c468cc8363d500db52c2ead036020723024a109d37346efaa761"},
{file = "h11-0.14.0.tar.gz", hash = "sha256:8f19fbbe99e72420ff35c00b27a34cb9937e902a8b810e2c88300c6f0a3b699d"},
@@ -858,6 +889,7 @@ version = "4.2.0"
description = "Pure-Python HTTP/2 protocol implementation"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "h2-4.2.0-py3-none-any.whl", hash = "sha256:479a53ad425bb29af087f3458a61d30780bc818e4ebcf01f0b536ba916462ed0"},
{file = "h2-4.2.0.tar.gz", hash = "sha256:c8a52129695e88b1a0578d8d2cc6842bbd79128ac685463b887ee278126ad01f"},
@@ -873,6 +905,7 @@ version = "4.1.0"
description = "Pure-Python HPACK header encoding"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "hpack-4.1.0-py3-none-any.whl", hash = "sha256:157ac792668d995c657d93111f46b4535ed114f0c9c8d672271bbec7eae1b496"},
{file = "hpack-4.1.0.tar.gz", hash = "sha256:ec5eca154f7056aa06f196a557655c5b009b382873ac8d1e66e79e87535f1dca"},
@@ -884,6 +917,7 @@ version = "1.0.7"
description = "A minimal low-level HTTP client."
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "httpcore-1.0.7-py3-none-any.whl", hash = "sha256:a3fff8f43dc260d5bd363d9f9cf1830fa3a458b332856f34282de498ed420edd"},
{file = "httpcore-1.0.7.tar.gz", hash = "sha256:8551cb62a169ec7162ac7be8d4817d561f60e08eaa485234898414bb5a8a0b4c"},
@@ -905,6 +939,7 @@ version = "0.28.1"
description = "The next generation HTTP client."
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "httpx-0.28.1-py3-none-any.whl", hash = "sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad"},
{file = "httpx-0.28.1.tar.gz", hash = "sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc"},
@@ -930,6 +965,7 @@ version = "6.1.0"
description = "Pure-Python HTTP/2 framing"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "hyperframe-6.1.0-py3-none-any.whl", hash = "sha256:b03380493a519fce58ea5af42e4a42317bf9bd425596f7a0835ffce80f1a42e5"},
{file = "hyperframe-6.1.0.tar.gz", hash = "sha256:f630908a00854a7adeabd6382b43923a4c4cd4b821fcb527e6ab9e15382a3b08"},
@@ -941,6 +977,7 @@ version = "3.10"
description = "Internationalized Domain Names in Applications (IDNA)"
optional = false
python-versions = ">=3.6"
groups = ["main"]
files = [
{file = "idna-3.10-py3-none-any.whl", hash = "sha256:946d195a0d259cbba61165e88e65941f16e9b36ea6ddb97f00452bae8b1287d3"},
{file = "idna-3.10.tar.gz", hash = "sha256:12f65c9b470abda6dc35cf8e63cc574b1c52b11df2c86030af0ac09b01b13ea9"},
@@ -955,6 +992,7 @@ version = "2.0.0"
description = "brain-dead simple config-ini parsing"
optional = false
python-versions = ">=3.7"
groups = ["dev", "test"]
files = [
{file = "iniconfig-2.0.0-py3-none-any.whl", hash = "sha256:b6a85871a79d2e3b22d2d1b94ac2824226a63c6b741c88f7ae975f18b6778374"},
{file = "iniconfig-2.0.0.tar.gz", hash = "sha256:2d91e135bf72d31a410b17c16da610a82cb55f6b0477d1a902134b24a455b8b3"},
@@ -978,6 +1016,7 @@ version = "5.13.2"
description = "A Python utility / library to sort Python imports."
optional = false
python-versions = ">=3.8.0"
groups = ["dev"]
files = [
{file = "isort-5.13.2-py3-none-any.whl", hash = "sha256:8ca5e72a8d85860d5a3fa69b8745237f2939afe12dbf656afbcb47fe72d947a6"},
{file = "isort-5.13.2.tar.gz", hash = "sha256:48fdfcb9face5d58a4f6dde2e72a1fb8dcaf8ab26f95ab49fab84c2ddefb0109"},
@@ -992,6 +1031,7 @@ version = "0.8.2"
description = "Fast iterable JSON parser."
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "jiter-0.8.2-cp310-cp310-macosx_10_12_x86_64.whl", hash = "sha256:ca8577f6a413abe29b079bc30f907894d7eb07a865c4df69475e868d73e71c7b"},
{file = "jiter-0.8.2-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:b25bd626bde7fb51534190c7e3cb97cee89ee76b76d7585580e22f34f5e3f393"},
@@ -1077,6 +1117,8 @@ version = "0.30.3"
description = "A package to repair broken json strings"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "json_repair-0.30.3-py3-none-any.whl", hash = "sha256:63bb588162b0958ae93d85356ecbe54c06b8c33f8a4834f93fa2719ea669804e"},
{file = "json_repair-0.30.3.tar.gz", hash = "sha256:0ac56e7ae9253ee9c507a7e1a3a26799c9b0bbe5e2bec1b2cc5053e90d5b05e3"},
@@ -1088,6 +1130,8 @@ version = "1.33"
description = "Apply JSON-Patches (RFC 6902)"
optional = false
python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*, !=3.4.*, !=3.5.*, !=3.6.*"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "jsonpatch-1.33-py2.py3-none-any.whl", hash = "sha256:0ae28c0cd062bbd8b8ecc26d7d164fbbea9652a1a3693f3b956c1eae5145dade"},
{file = "jsonpatch-1.33.tar.gz", hash = "sha256:9fcd4009c41e6d12348b4a0ff2563ba56a2923a7dfee731d004e212e1ee5030c"},
@@ -1102,6 +1146,8 @@ version = "3.0.0"
description = "Identify specific nodes in a JSON document (RFC 6901)"
optional = false
python-versions = ">=3.7"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "jsonpointer-3.0.0-py2.py3-none-any.whl", hash = "sha256:13e088adc14fca8b6aa8177c044e12701e6ad4b28ff10e65f2267a90109c9942"},
{file = "jsonpointer-3.0.0.tar.gz", hash = "sha256:2b2d729f2091522d61c3b31f82e11870f60b68f43fbc705cb76bf4b832af59ef"},
@@ -1113,6 +1159,8 @@ version = "0.3.19"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "langchain-0.3.19-py3-none-any.whl", hash = "sha256:1e16d97db9106640b7de4c69f8f5ed22eeda56b45b9241279e83f111640eff16"},
{file = "langchain-0.3.19.tar.gz", hash = "sha256:b96f8a445f01d15d522129ffe77cc89c8468dbd65830d153a676de8f6b899e7b"},
@@ -1151,12 +1199,54 @@ openai = ["langchain-openai"]
together = ["langchain-together"]
xai = ["langchain-xai"]
[[package]]
name = "langchain"
version = "0.3.20"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "python_version >= \"3.10\" and python_version < \"3.12\" and extra == \"graph\""
files = [
{file = "langchain-0.3.20-py3-none-any.whl", hash = "sha256:273287f8e61ffdf7e811cf8799e6a71e9381325b8625fd6618900faba79cfdd0"},
{file = "langchain-0.3.20.tar.gz", hash = "sha256:edcc3241703e1f6557ef5a5c35cd56f9ccc25ff12e38b4829c66d94971737a93"},
]
[package.dependencies]
async-timeout = {version = ">=4.0.0,<5.0.0", markers = "python_version < \"3.11\""}
langchain-core = ">=0.3.41,<1.0.0"
langchain-text-splitters = ">=0.3.6,<1.0.0"
langsmith = ">=0.1.17,<0.4"
pydantic = ">=2.7.4,<3.0.0"
PyYAML = ">=5.3"
requests = ">=2,<3"
SQLAlchemy = ">=1.4,<3"
[package.extras]
anthropic = ["langchain-anthropic"]
aws = ["langchain-aws"]
cohere = ["langchain-cohere"]
community = ["langchain-community"]
deepseek = ["langchain-deepseek"]
fireworks = ["langchain-fireworks"]
google-genai = ["langchain-google-genai"]
google-vertexai = ["langchain-google-vertexai"]
groq = ["langchain-groq"]
huggingface = ["langchain-huggingface"]
mistralai = ["langchain-mistralai"]
ollama = ["langchain-ollama"]
openai = ["langchain-openai"]
together = ["langchain-together"]
xai = ["langchain-xai"]
[[package]]
name = "langchain-core"
version = "0.3.40"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "langchain_core-0.3.40-py3-none-any.whl", hash = "sha256:9f31358741f10a13db8531e8288b8a5ae91904018c5c2e6f739d6645a98fca03"},
{file = "langchain_core-0.3.40.tar.gz", hash = "sha256:893a238b38491967c804662c1ec7c3e6ebaf223d1125331249c3cf3862ff2746"},
@@ -1174,12 +1264,36 @@ PyYAML = ">=5.3"
tenacity = ">=8.1.0,<8.4.0 || >8.4.0,<10.0.0"
typing-extensions = ">=4.7"
[[package]]
name = "langchain-core"
version = "0.3.45"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "python_version >= \"3.10\" and python_version < \"3.12\" and extra == \"graph\""
files = [
{file = "langchain_core-0.3.45-py3-none-any.whl", hash = "sha256:fe560d644c102c3f5dcfb44eb5295e26d22deab259fdd084f6b1b55a0350b77c"},
{file = "langchain_core-0.3.45.tar.gz", hash = "sha256:a39b8446495d1ea97311aa726478c0a13ef1d77cb7644350bad6d9d3c0141a0c"},
]
[package.dependencies]
jsonpatch = ">=1.33,<2.0"
langsmith = ">=0.1.125,<0.4"
packaging = ">=23.2,<25"
pydantic = {version = ">=2.5.2,<3.0.0", markers = "python_full_version < \"3.12.4\""}
PyYAML = ">=5.3"
tenacity = ">=8.1.0,<8.4.0 || >8.4.0,<10.0.0"
typing-extensions = ">=4.7"
[[package]]
name = "langchain-neo4j"
version = "0.4.0"
description = "An integration package connecting Neo4j and LangChain"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "langchain_neo4j-0.4.0-py3-none-any.whl", hash = "sha256:2760b5757e7a402884cf3419830217651df97fe4f44b3fec6c96b14b6d7fd18e"},
{file = "langchain_neo4j-0.4.0.tar.gz", hash = "sha256:3f059a66411cec1062a2b8c44953a70d0fff9e123e9fb1d6b3f17a0bef6d6114"},
@@ -1197,6 +1311,8 @@ version = "0.3.6"
description = "LangChain text splitting utilities"
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "langchain_text_splitters-0.3.6-py3-none-any.whl", hash = "sha256:e5d7b850f6c14259ea930be4a964a65fa95d9df7e1dbdd8bad8416db72292f4e"},
{file = "langchain_text_splitters-0.3.6.tar.gz", hash = "sha256:c537972f4b7c07451df431353a538019ad9dadff7a1073ea363946cea97e1bee"},
@@ -1211,6 +1327,8 @@ version = "0.3.11"
description = "Client library to connect to the LangSmith LLM Tracing and Evaluation Platform."
optional = false
python-versions = "<4.0,>=3.9"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "langsmith-0.3.11-py3-none-any.whl", hash = "sha256:0cca22737ef07d3b038a437c141deda37e00add56022582680188b681bec095e"},
{file = "langsmith-0.3.11.tar.gz", hash = "sha256:ddf29d24352e99de79c9618aaf95679214324e146c5d3d9475a7ddd2870018b1"},
@@ -1238,6 +1356,7 @@ version = "1.6"
description = "An implementation of time.monotonic() for Python 2 & < 3.3"
optional = false
python-versions = "*"
groups = ["main"]
files = [
{file = "monotonic-1.6-py2.py3-none-any.whl", hash = "sha256:68687e19a14f11f26d140dd5c86f3dba4bf5df58003000ed467e0e2a69bca96c"},
{file = "monotonic-1.6.tar.gz", hash = "sha256:3a55207bcfed53ddd5c5bae174524062935efed17792e9de2ad0205ce9ad63f7"},
@@ -1249,6 +1368,8 @@ version = "6.1.0"
description = "multidict implementation"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "multidict-6.1.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:3380252550e372e8511d49481bd836264c009adb826b23fefcc5dd3c69692f60"},
{file = "multidict-6.1.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:99f826cbf970077383d7de805c0681799491cb939c25450b9b5b3ced03ca99f1"},
@@ -1353,6 +1474,8 @@ version = "5.28.1"
description = "Neo4j Bolt driver for Python"
optional = false
python-versions = ">=3.7"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "neo4j-5.28.1-py3-none-any.whl", hash = "sha256:6755ef9e5f4e14b403aef1138fb6315b120631a0075c138b5ddb2a06b87b09fd"},
{file = "neo4j-5.28.1.tar.gz", hash = "sha256:ae8e37a1d895099062c75bc359b2cce62099baac7be768d0eba7180c1298e214"},
@@ -1372,6 +1495,8 @@ version = "1.5.0"
description = "Python package to allow easy integration to Neo4j's GraphRAG features"
optional = false
python-versions = "<4.0.0,>=3.9.0"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "neo4j_graphrag-1.5.0-py3-none-any.whl", hash = "sha256:d778be8476aa758ff10043b373c287ada0abe90bcfb9c34e8cb4d2c4bbd20240"},
{file = "neo4j_graphrag-1.5.0.tar.gz", hash = "sha256:5367c08f128e83ca7028f2112f6e0b69541400adee40cdaf6e30c3563036acb3"},
@@ -1389,9 +1514,9 @@ types-pyyaml = ">=6.0.12.20240917,<7.0.0.0"
[package.extras]
anthropic = ["anthropic (>=0.36.0,<0.37.0)"]
cohere = ["cohere (>=5.9.0,<6.0.0)"]
experimental = ["langchain-text-splitters (>=0.3.0,<0.4.0)", "llama-index (>=0.10.55,<0.11.0)", "pygraphviz (>=1.0.0,<2.0.0)", "pygraphviz (>=1.13.0,<2.0.0)"]
experimental = ["langchain-text-splitters (>=0.3.0,<0.4.0)", "llama-index (>=0.10.55,<0.11.0)", "pygraphviz (>=1.0.0,<2.0.0) ; python_version < \"3.10\"", "pygraphviz (>=1.13.0,<2.0.0) ; python_version >= \"3.10\" and python_full_version < \"4.0.0\""]
google = ["google-cloud-aiplatform (>=1.66.0,<2.0.0)"]
kg-creation-tools = ["pygraphviz (>=1.0.0,<2.0.0)", "pygraphviz (>=1.13.0,<2.0.0)"]
kg-creation-tools = ["pygraphviz (>=1.0.0,<2.0.0) ; python_version < \"3.10\"", "pygraphviz (>=1.13.0,<2.0.0) ; python_version >= \"3.10\" and python_full_version < \"4.0.0\""]
mistralai = ["mistralai (>=1.0.3,<2.0.0)"]
ollama = ["ollama (>=0.4.4,<0.5.0)"]
openai = ["openai (>=1.51.1,<2.0.0)"]
@@ -1406,6 +1531,8 @@ version = "1.26.4"
description = "Fundamental package for array computing in Python"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "python_version < \"3.10\""
files = [
{file = "numpy-1.26.4-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:9ff0f4f29c51e2803569d7a51c2304de5554655a60c5d776e35b4a41413830d0"},
{file = "numpy-1.26.4-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:2e4ee3380d6de9c9ec04745830fd9e2eccb3e6cf790d39d7b98ffd19b0dd754a"},
@@ -1445,12 +1572,79 @@ files = [
{file = "numpy-1.26.4.tar.gz", hash = "sha256:2a02aba9ed12e4ac4eb3ea9421c420301a0c6460d9830d74a9df87efa4912010"},
]
[[package]]
name = "numpy"
version = "2.2.4"
description = "Fundamental package for array computing in Python"
optional = false
python-versions = ">=3.10"
groups = ["main"]
markers = "python_version >= \"3.10\""
files = [
{file = "numpy-2.2.4-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:8146f3550d627252269ac42ae660281d673eb6f8b32f113538e0cc2a9aed42b9"},
{file = "numpy-2.2.4-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:e642d86b8f956098b564a45e6f6ce68a22c2c97a04f5acd3f221f57b8cb850ae"},
{file = "numpy-2.2.4-cp310-cp310-macosx_14_0_arm64.whl", hash = "sha256:a84eda42bd12edc36eb5b53bbcc9b406820d3353f1994b6cfe453a33ff101775"},
{file = "numpy-2.2.4-cp310-cp310-macosx_14_0_x86_64.whl", hash = "sha256:4ba5054787e89c59c593a4169830ab362ac2bee8a969249dc56e5d7d20ff8df9"},
{file = "numpy-2.2.4-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7716e4a9b7af82c06a2543c53ca476fa0b57e4d760481273e09da04b74ee6ee2"},
{file = "numpy-2.2.4-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:adf8c1d66f432ce577d0197dceaac2ac00c0759f573f28516246351c58a85020"},
{file = "numpy-2.2.4-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:218f061d2faa73621fa23d6359442b0fc658d5b9a70801373625d958259eaca3"},
{file = "numpy-2.2.4-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:df2f57871a96bbc1b69733cd4c51dc33bea66146b8c63cacbfed73eec0883017"},
{file = "numpy-2.2.4-cp310-cp310-win32.whl", hash = "sha256:a0258ad1f44f138b791327961caedffbf9612bfa504ab9597157806faa95194a"},
{file = "numpy-2.2.4-cp310-cp310-win_amd64.whl", hash = "sha256:0d54974f9cf14acf49c60f0f7f4084b6579d24d439453d5fc5805d46a165b542"},
{file = "numpy-2.2.4-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:e9e0a277bb2eb5d8a7407e14688b85fd8ad628ee4e0c7930415687b6564207a4"},
{file = "numpy-2.2.4-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:9eeea959168ea555e556b8188da5fa7831e21d91ce031e95ce23747b7609f8a4"},
{file = "numpy-2.2.4-cp311-cp311-macosx_14_0_arm64.whl", hash = "sha256:bd3ad3b0a40e713fc68f99ecfd07124195333f1e689387c180813f0e94309d6f"},
{file = "numpy-2.2.4-cp311-cp311-macosx_14_0_x86_64.whl", hash = "sha256:cf28633d64294969c019c6df4ff37f5698e8326db68cc2b66576a51fad634880"},
{file = "numpy-2.2.4-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:2fa8fa7697ad1646b5c93de1719965844e004fcad23c91228aca1cf0800044a1"},
{file = "numpy-2.2.4-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:f4162988a360a29af158aeb4a2f4f09ffed6a969c9776f8f3bdee9b06a8ab7e5"},
{file = "numpy-2.2.4-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:892c10d6a73e0f14935c31229e03325a7b3093fafd6ce0af704be7f894d95687"},
{file = "numpy-2.2.4-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:db1f1c22173ac1c58db249ae48aa7ead29f534b9a948bc56828337aa84a32ed6"},
{file = "numpy-2.2.4-cp311-cp311-win32.whl", hash = "sha256:ea2bb7e2ae9e37d96835b3576a4fa4b3a97592fbea8ef7c3587078b0068b8f09"},
{file = "numpy-2.2.4-cp311-cp311-win_amd64.whl", hash = "sha256:f7de08cbe5551911886d1ab60de58448c6df0f67d9feb7d1fb21e9875ef95e91"},
{file = "numpy-2.2.4-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:a7b9084668aa0f64e64bd00d27ba5146ef1c3a8835f3bd912e7a9e01326804c4"},
{file = "numpy-2.2.4-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:dbe512c511956b893d2dacd007d955a3f03d555ae05cfa3ff1c1ff6df8851854"},
{file = "numpy-2.2.4-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:bb649f8b207ab07caebba230d851b579a3c8711a851d29efe15008e31bb4de24"},
{file = "numpy-2.2.4-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:f34dc300df798742b3d06515aa2a0aee20941c13579d7a2f2e10af01ae4901ee"},
{file = "numpy-2.2.4-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:c3f7ac96b16955634e223b579a3e5798df59007ca43e8d451a0e6a50f6bfdfba"},
{file = "numpy-2.2.4-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:4f92084defa704deadd4e0a5ab1dc52d8ac9e8a8ef617f3fbb853e79b0ea3592"},
{file = "numpy-2.2.4-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:7a4e84a6283b36632e2a5b56e121961f6542ab886bc9e12f8f9818b3c266bfbb"},
{file = "numpy-2.2.4-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:11c43995255eb4127115956495f43e9343736edb7fcdb0d973defd9de14cd84f"},
{file = "numpy-2.2.4-cp312-cp312-win32.whl", hash = "sha256:65ef3468b53269eb5fdb3a5c09508c032b793da03251d5f8722b1194f1790c00"},
{file = "numpy-2.2.4-cp312-cp312-win_amd64.whl", hash = "sha256:2aad3c17ed2ff455b8eaafe06bcdae0062a1db77cb99f4b9cbb5f4ecb13c5146"},
{file = "numpy-2.2.4-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:1cf4e5c6a278d620dee9ddeb487dc6a860f9b199eadeecc567f777daace1e9e7"},
{file = "numpy-2.2.4-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:1974afec0b479e50438fc3648974268f972e2d908ddb6d7fb634598cdb8260a0"},
{file = "numpy-2.2.4-cp313-cp313-macosx_14_0_arm64.whl", hash = "sha256:79bd5f0a02aa16808fcbc79a9a376a147cc1045f7dfe44c6e7d53fa8b8a79392"},
{file = "numpy-2.2.4-cp313-cp313-macosx_14_0_x86_64.whl", hash = "sha256:3387dd7232804b341165cedcb90694565a6015433ee076c6754775e85d86f1fc"},
{file = "numpy-2.2.4-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:6f527d8fdb0286fd2fd97a2a96c6be17ba4232da346931d967a0630050dfd298"},
{file = "numpy-2.2.4-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:bce43e386c16898b91e162e5baaad90c4b06f9dcbe36282490032cec98dc8ae7"},
{file = "numpy-2.2.4-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:31504f970f563d99f71a3512d0c01a645b692b12a63630d6aafa0939e52361e6"},
{file = "numpy-2.2.4-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:81413336ef121a6ba746892fad881a83351ee3e1e4011f52e97fba79233611fd"},
{file = "numpy-2.2.4-cp313-cp313-win32.whl", hash = "sha256:f486038e44caa08dbd97275a9a35a283a8f1d2f0ee60ac260a1790e76660833c"},
{file = "numpy-2.2.4-cp313-cp313-win_amd64.whl", hash = "sha256:207a2b8441cc8b6a2a78c9ddc64d00d20c303d79fba08c577752f080c4007ee3"},
{file = "numpy-2.2.4-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:8120575cb4882318c791f839a4fd66161a6fa46f3f0a5e613071aae35b5dd8f8"},
{file = "numpy-2.2.4-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:a761ba0fa886a7bb33c6c8f6f20213735cb19642c580a931c625ee377ee8bd39"},
{file = "numpy-2.2.4-cp313-cp313t-macosx_14_0_arm64.whl", hash = "sha256:ac0280f1ba4a4bfff363a99a6aceed4f8e123f8a9b234c89140f5e894e452ecd"},
{file = "numpy-2.2.4-cp313-cp313t-macosx_14_0_x86_64.whl", hash = "sha256:879cf3a9a2b53a4672a168c21375166171bc3932b7e21f622201811c43cdd3b0"},
{file = "numpy-2.2.4-cp313-cp313t-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:f05d4198c1bacc9124018109c5fba2f3201dbe7ab6e92ff100494f236209c960"},
{file = "numpy-2.2.4-cp313-cp313t-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e2f085ce2e813a50dfd0e01fbfc0c12bbe5d2063d99f8b29da30e544fb6483b8"},
{file = "numpy-2.2.4-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:92bda934a791c01d6d9d8e038363c50918ef7c40601552a58ac84c9613a665bc"},
{file = "numpy-2.2.4-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:ee4d528022f4c5ff67332469e10efe06a267e32f4067dc76bb7e2cddf3cd25ff"},
{file = "numpy-2.2.4-cp313-cp313t-win32.whl", hash = "sha256:05c076d531e9998e7e694c36e8b349969c56eadd2cdcd07242958489d79a7286"},
{file = "numpy-2.2.4-cp313-cp313t-win_amd64.whl", hash = "sha256:188dcbca89834cc2e14eb2f106c96d6d46f200fe0200310fc29089657379c58d"},
{file = "numpy-2.2.4-pp310-pypy310_pp73-macosx_10_15_x86_64.whl", hash = "sha256:7051ee569db5fbac144335e0f3b9c2337e0c8d5c9fee015f259a5bd70772b7e8"},
{file = "numpy-2.2.4-pp310-pypy310_pp73-macosx_14_0_x86_64.whl", hash = "sha256:ab2939cd5bec30a7430cbdb2287b63151b77cf9624de0532d629c9a1c59b1d5c"},
{file = "numpy-2.2.4-pp310-pypy310_pp73-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:d0f35b19894a9e08639fd60a1ec1978cb7f5f7f1eace62f38dd36be8aecdef4d"},
{file = "numpy-2.2.4-pp310-pypy310_pp73-win_amd64.whl", hash = "sha256:b4adfbbc64014976d2f91084915ca4e626fbf2057fb81af209c1a6d776d23e3d"},
{file = "numpy-2.2.4.tar.gz", hash = "sha256:9ba03692a45d3eef66559efe1d1096c4b9b75c0986b5dff5530c378fb8331d4f"},
]
[[package]]
name = "openai"
version = "1.65.2"
description = "The official Python library for the openai API"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "openai-1.65.2-py3-none-any.whl", hash = "sha256:27d9fe8de876e31394c2553c4e6226378b6ed85e480f586ccfe25b7193fb1750"},
{file = "openai-1.65.2.tar.gz", hash = "sha256:729623efc3fd91c956f35dd387fa5c718edd528c4bed9f00b40ef290200fb2ce"},
@@ -1476,6 +1670,8 @@ version = "3.10.15"
description = "Fast, correct Python JSON library supporting dataclasses, datetimes, and numpy"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\" and platform_python_implementation != \"PyPy\""
files = [
{file = "orjson-3.10.15-cp310-cp310-macosx_10_15_x86_64.macosx_11_0_arm64.macosx_10_15_universal2.whl", hash = "sha256:552c883d03ad185f720d0c09583ebde257e41b9521b74ff40e08b7dec4559c04"},
{file = "orjson-3.10.15-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:616e3e8d438d02e4854f70bfdc03a6bcdb697358dbaa6bcd19cbe24d24ece1f8"},
@@ -1564,10 +1760,12 @@ version = "24.2"
description = "Core utilities for Python packages"
optional = false
python-versions = ">=3.8"
groups = ["main", "dev", "test"]
files = [
{file = "packaging-24.2-py3-none-any.whl", hash = "sha256:09abb1bccd265c01f4a3aa3f7a7db064b36514d2cba19a2f694fe6150451a759"},
{file = "packaging-24.2.tar.gz", hash = "sha256:c228a6dc5e932d346bc5739379109d49e8853dd8223571c7c5b55260edc0b97f"},
]
markers = {main = "extra == \"graph\""}
[[package]]
name = "pluggy"
@@ -1575,6 +1773,7 @@ version = "1.5.0"
description = "plugin and hook calling mechanisms for python"
optional = false
python-versions = ">=3.8"
groups = ["dev", "test"]
files = [
{file = "pluggy-1.5.0-py3-none-any.whl", hash = "sha256:44e1ad92c8ca002de6377e165f3e0f1be63266ab4d554740532335b9d75ea669"},
{file = "pluggy-1.5.0.tar.gz", hash = "sha256:2cffa88e94fdc978c4c574f15f9e59b7f4201d439195c3715ca9e2486f1d0cf1"},
@@ -1590,6 +1789,7 @@ version = "2.10.1"
description = "Wraps the portalocker recipe for easy usage"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "portalocker-2.10.1-py3-none-any.whl", hash = "sha256:53a5984ebc86a025552264b459b46a2086e269b21823cb572f8f28ee759e45bf"},
{file = "portalocker-2.10.1.tar.gz", hash = "sha256:ef1bf844e878ab08aee7e40184156e1151f228f103aa5c6bd0724cc330960f8f"},
@@ -1609,6 +1809,7 @@ version = "3.18.0"
description = "Integrate PostHog into any python application."
optional = false
python-versions = "*"
groups = ["main"]
files = [
{file = "posthog-3.18.0-py2.py3-none-any.whl", hash = "sha256:88f93cc670158ea7a569629d4def77c3b714489a52b0a818b8dd6103669c652d"},
{file = "posthog-3.18.0.tar.gz", hash = "sha256:8882fe1ee6763bafab0d82a506f52d752a8031c90eea17b4db4331fb33cf1092"},
@@ -1634,6 +1835,8 @@ version = "0.3.0"
description = "Accelerated property cache"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "propcache-0.3.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:efa44f64c37cc30c9f05932c740a8b40ce359f51882c70883cc95feac842da4d"},
{file = "propcache-0.3.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:2383a17385d9800b6eb5855c2f05ee550f803878f344f58b6e194de08b96352c"},
@@ -1741,6 +1944,7 @@ version = "5.29.3"
description = ""
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "protobuf-5.29.3-cp310-abi3-win32.whl", hash = "sha256:3ea51771449e1035f26069c4c7fd51fba990d07bc55ba80701c78f886bf9c888"},
{file = "protobuf-5.29.3-cp310-abi3-win_amd64.whl", hash = "sha256:a4fa6f80816a9a0678429e84973f2f98cbc218cca434abe8db2ad0bffc98503a"},
@@ -1761,6 +1965,7 @@ version = "2.9.10"
description = "psycopg2 - Python-PostgreSQL Database Adapter"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "psycopg2-binary-2.9.10.tar.gz", hash = "sha256:4b3df0e6990aa98acda57d983942eff13d824135fe2250e6522edaa782a06de2"},
{file = "psycopg2_binary-2.9.10-cp310-cp310-macosx_12_0_x86_64.whl", hash = "sha256:0ea8e3d0ae83564f2fc554955d327fa081d065c8ca5cc6d2abb643e2c9c1200f"},
@@ -1838,6 +2043,8 @@ version = "2.22"
description = "C parser in Python"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\" and platform_python_implementation == \"PyPy\""
files = [
{file = "pycparser-2.22-py3-none-any.whl", hash = "sha256:c3702b6d3dd8c7abc1afa565d7e63d53a1d0bd86cdc24edd75470f4de499cfcc"},
{file = "pycparser-2.22.tar.gz", hash = "sha256:491c8be9c040f5390f5bf44a5b07752bd07f56edf992381b05c701439eec10f6"},
@@ -1849,6 +2056,7 @@ version = "2.10.6"
description = "Data validation using Python type hints"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "pydantic-2.10.6-py3-none-any.whl", hash = "sha256:427d664bf0b8a2b34ff5dd0f5a18df00591adcee7198fbd71981054cef37b584"},
{file = "pydantic-2.10.6.tar.gz", hash = "sha256:ca5daa827cce33de7a42be142548b0096bf05a7e7b365aebfa5f8eeec7128236"},
@@ -1869,6 +2077,7 @@ version = "2.27.2"
description = "Core functionality for Pydantic validation and serialization"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "pydantic_core-2.27.2-cp310-cp310-macosx_10_12_x86_64.whl", hash = "sha256:2d367ca20b2f14095a8f4fa1210f5a7b78b8a20009ecced6b12818f455b1e9fa"},
{file = "pydantic_core-2.27.2-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:491a2b73db93fab69731eaee494f320faa4e093dbed776be1a829c2eb222c34c"},
@@ -1981,6 +2190,8 @@ version = "4.3.1"
description = "A pure-python PDF library capable of splitting, merging, cropping, and transforming PDF files"
optional = false
python-versions = ">=3.6"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "pypdf-4.3.1-py3-none-any.whl", hash = "sha256:64b31da97eda0771ef22edb1bfecd5deee4b72c3d1736b7df2689805076d6418"},
{file = "pypdf-4.3.1.tar.gz", hash = "sha256:b2f37fe9a3030aa97ca86067a56ba3f9d3565f9a791b305c7355d8392c30d91b"},
@@ -1990,10 +2201,10 @@ files = [
typing_extensions = {version = ">=4.0", markers = "python_version < \"3.11\""}
[package.extras]
crypto = ["PyCryptodome", "cryptography"]
crypto = ["PyCryptodome ; python_version == \"3.6\"", "cryptography ; python_version >= \"3.7\""]
dev = ["black", "flit", "pip-tools", "pre-commit (<2.18.0)", "pytest-cov", "pytest-socket", "pytest-timeout", "pytest-xdist", "wheel"]
docs = ["myst_parser", "sphinx", "sphinx_rtd_theme"]
full = ["Pillow (>=8.0.0)", "PyCryptodome", "cryptography"]
full = ["Pillow (>=8.0.0)", "PyCryptodome ; python_version == \"3.6\"", "cryptography ; python_version >= \"3.7\""]
image = ["Pillow (>=8.0.0)"]
[[package]]
@@ -2002,6 +2213,7 @@ version = "8.3.5"
description = "pytest: simple powerful testing with Python"
optional = false
python-versions = ">=3.8"
groups = ["dev", "test"]
files = [
{file = "pytest-8.3.5-py3-none-any.whl", hash = "sha256:c69214aa47deac29fad6c2a4f590b9c4a9fdb16a403176fe154b79c0b4d4d820"},
{file = "pytest-8.3.5.tar.gz", hash = "sha256:f4efe70cc14e511565ac476b57c279e12a855b11f48f212af1080ef2263d3845"},
@@ -2024,6 +2236,7 @@ version = "2.9.0.post0"
description = "Extensions to the standard Python datetime module"
optional = false
python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,>=2.7"
groups = ["main"]
files = [
{file = "python-dateutil-2.9.0.post0.tar.gz", hash = "sha256:37dd54208da7e1cd875388217d5e00ebd4179249f90fb72437e91a35459a0ad3"},
{file = "python_dateutil-2.9.0.post0-py2.py3-none-any.whl", hash = "sha256:a8b2bc7bffae282281c8140a97d3aa9c14da0b136dfe83f850eea9a5f7470427"},
@@ -2038,6 +2251,7 @@ version = "2024.2"
description = "World timezone definitions, modern and historical"
optional = false
python-versions = "*"
groups = ["main"]
files = [
{file = "pytz-2024.2-py2.py3-none-any.whl", hash = "sha256:31c7c1817eb7fae7ca4b8c7ee50c72f93aa2dd863de768e1ef4245d426aa0725"},
{file = "pytz-2024.2.tar.gz", hash = "sha256:2aa355083c50a0f93fa581709deac0c9ad65cca8a9e9beac660adcbd493c798a"},
@@ -2049,6 +2263,8 @@ version = "308"
description = "Python for Window Extensions"
optional = false
python-versions = "*"
groups = ["main"]
markers = "platform_system == \"Windows\""
files = [
{file = "pywin32-308-cp310-cp310-win32.whl", hash = "sha256:796ff4426437896550d2981b9c2ac0ffd75238ad9ea2d3bfa67a1abd546d262e"},
{file = "pywin32-308-cp310-cp310-win_amd64.whl", hash = "sha256:4fc888c59b3c0bef905ce7eb7e2106a07712015ea1c8234b703a088d46110e8e"},
@@ -2076,6 +2292,8 @@ version = "6.0.2"
description = "YAML parser and emitter for Python"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "PyYAML-6.0.2-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:0a9a2848a5b7feac301353437eb7d5957887edbf81d56e903999a75a3d743086"},
{file = "PyYAML-6.0.2-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:29717114e51c84ddfba879543fb232a6ed60086602313ca38cce623c1d62cfbf"},
@@ -2132,36 +2350,13 @@ files = [
{file = "pyyaml-6.0.2.tar.gz", hash = "sha256:d584d9ec91ad65861cc08d42e834324ef890a082e591037abe114850ff7bbc3e"},
]
[[package]]
name = "qdrant-client"
version = "1.12.1"
description = "Client library for the Qdrant vector search engine"
optional = false
python-versions = ">=3.8"
files = [
{file = "qdrant_client-1.12.1-py3-none-any.whl", hash = "sha256:b2d17ce18e9e767471368380dd3bbc4a0e3a0e2061fedc9af3542084b48451e0"},
{file = "qdrant_client-1.12.1.tar.gz", hash = "sha256:35e8e646f75b7b883b3d2d0ee4c69c5301000bba41c82aa546e985db0f1aeb72"},
]
[package.dependencies]
grpcio = ">=1.41.0"
grpcio-tools = ">=1.41.0"
httpx = {version = ">=0.20.0", extras = ["http2"]}
numpy = {version = ">=1.26", markers = "python_version >= \"3.12\""}
portalocker = ">=2.7.0,<3.0.0"
pydantic = ">=1.10.8"
urllib3 = ">=1.26.14,<3"
[package.extras]
fastembed = ["fastembed (==0.3.6)"]
fastembed-gpu = ["fastembed-gpu (==0.3.6)"]
[[package]]
name = "qdrant-client"
version = "1.13.2"
description = "Client library for the Qdrant vector search engine"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "qdrant_client-1.13.2-py3-none-any.whl", hash = "sha256:db97e759bd3f8d483a383984ba4c2a158eef56f2188d83df7771591d43de2201"},
{file = "qdrant_client-1.13.2.tar.gz", hash = "sha256:c8cce87ce67b006f49430a050a35c85b78e3b896c0c756dafc13bdeca543ec13"},
@@ -2174,7 +2369,8 @@ httpx = {version = ">=0.20.0", extras = ["http2"]}
numpy = [
{version = ">=1.21", markers = "python_version >= \"3.10\" and python_version < \"3.12\""},
{version = ">=1.21,<2.1.0", markers = "python_version < \"3.10\""},
{version = ">=1.26", markers = "python_version >= \"3.12\" and python_version < \"3.13\""},
{version = ">=1.26", markers = "python_version == \"3.12\""},
{version = ">=2.1.0", markers = "python_version >= \"3.13\""},
]
portalocker = ">=2.7.0,<3.0.0"
pydantic = ">=1.10.8"
@@ -2190,6 +2386,8 @@ version = "0.2.2"
description = "Various BM25 algorithms for document ranking"
optional = false
python-versions = "*"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "rank_bm25-0.2.2-py3-none-any.whl", hash = "sha256:7bd4a95571adadfc271746fa146a4bcfd89c0cf731e49c3d1ad863290adbe8ae"},
{file = "rank_bm25-0.2.2.tar.gz", hash = "sha256:096ccef76f8188563419aaf384a02f0ea459503fdf77901378d4fd9d87e5e51d"},
@@ -2207,6 +2405,7 @@ version = "2.32.3"
description = "Python HTTP for Humans."
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "requests-2.32.3-py3-none-any.whl", hash = "sha256:70761cfe03c773ceb22aa2f671b4757976145175cdfca038c02654d061d6dcc6"},
{file = "requests-2.32.3.tar.gz", hash = "sha256:55365417734eb18255590a9ff9eb97e9e1da868d4ccd6402399eaf68af20a760"},
@@ -2228,6 +2427,8 @@ version = "1.0.0"
description = "A utility belt for advanced users of python-requests"
optional = false
python-versions = ">=2.7, !=3.0.*, !=3.1.*, !=3.2.*, !=3.3.*"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "requests-toolbelt-1.0.0.tar.gz", hash = "sha256:7681a0a3d047012b5bdc0ee37d7f8f07ebe76ab08caeccfc3921ce23c88d5bc6"},
{file = "requests_toolbelt-1.0.0-py2.py3-none-any.whl", hash = "sha256:cccfdd665f0a24fcf4726e690f65639d272bb0637b9b92dfd91a5568ccf6bd06"},
@@ -2242,6 +2443,7 @@ version = "0.6.9"
description = "An extremely fast Python linter and code formatter, written in Rust."
optional = false
python-versions = ">=3.7"
groups = ["dev"]
files = [
{file = "ruff-0.6.9-py3-none-linux_armv6l.whl", hash = "sha256:064df58d84ccc0ac0fcd63bc3090b251d90e2a372558c0f057c3f75ed73e1ccd"},
{file = "ruff-0.6.9-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:140d4b5c9f5fc7a7b074908a78ab8d384dd7f6510402267bc76c37195c02a7ec"},
@@ -2269,6 +2471,7 @@ version = "75.8.2"
description = "Easily download, build, install, upgrade, and uninstall Python packages"
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "setuptools-75.8.2-py3-none-any.whl", hash = "sha256:558e47c15f1811c1fa7adbd0096669bf76c1d3f433f58324df69f3f5ecac4e8f"},
{file = "setuptools-75.8.2.tar.gz", hash = "sha256:4880473a969e5f23f2a2be3646b2dfd84af9028716d398e46192f84bc36900d2"},
@@ -2289,6 +2492,7 @@ version = "1.17.0"
description = "Python 2 and 3 compatibility utilities"
optional = false
python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,>=2.7"
groups = ["main"]
files = [
{file = "six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274"},
{file = "six-1.17.0.tar.gz", hash = "sha256:ff70335d468e7eb6ec65b95b99d3a2836546063f63acc5171de367e834932a81"},
@@ -2300,6 +2504,7 @@ version = "1.3.1"
description = "Sniff out which async library your code is running under"
optional = false
python-versions = ">=3.7"
groups = ["main"]
files = [
{file = "sniffio-1.3.1-py3-none-any.whl", hash = "sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2"},
{file = "sniffio-1.3.1.tar.gz", hash = "sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc"},
@@ -2311,6 +2516,7 @@ version = "2.0.38"
description = "Database Abstraction Library"
optional = false
python-versions = ">=3.7"
groups = ["main"]
files = [
{file = "SQLAlchemy-2.0.38-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:5e1d9e429028ce04f187a9f522818386c8b076723cdbe9345708384f49ebcec6"},
{file = "SQLAlchemy-2.0.38-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:b87a90f14c68c925817423b0424381f0e16d80fc9a1a1046ef202ab25b19a444"},
@@ -2406,6 +2612,8 @@ version = "9.0.0"
description = "Retry code until it succeeds"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "tenacity-9.0.0-py3-none-any.whl", hash = "sha256:93de0c98785b27fcf659856aa9f54bfbd399e29969b0621bc7f762bd441b4539"},
{file = "tenacity-9.0.0.tar.gz", hash = "sha256:807f37ca97d62aa361264d497b0e31e92b8027044942bfa756160d908320d73b"},
@@ -2421,6 +2629,8 @@ version = "2.2.1"
description = "A lil' TOML parser"
optional = false
python-versions = ">=3.8"
groups = ["dev", "test"]
markers = "python_version < \"3.11\""
files = [
{file = "tomli-2.2.1-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:678e4fa69e4575eb77d103de3df8a895e1591b48e740211bd1067378c69e8249"},
{file = "tomli-2.2.1-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:023aa114dd824ade0100497eb2318602af309e5a55595f76b626d6d9f3b7b0a6"},
@@ -2462,6 +2672,7 @@ version = "4.67.1"
description = "Fast, Extensible Progress Meter"
optional = false
python-versions = ">=3.7"
groups = ["main"]
files = [
{file = "tqdm-4.67.1-py3-none-any.whl", hash = "sha256:26445eca388f82e72884e0d580d5464cd801a3ea01e63e5601bdff9ba6a48de2"},
{file = "tqdm-4.67.1.tar.gz", hash = "sha256:f8aef9c52c08c13a65f30ea34f4e5aac3fd1a34959879d7e59e63027286627f2"},
@@ -2483,6 +2694,8 @@ version = "6.0.12.20241230"
description = "Typing stubs for PyYAML"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "types_PyYAML-6.0.12.20241230-py3-none-any.whl", hash = "sha256:fa4d32565219b68e6dee5f67534c722e53c00d1cfc09c435ef04d7353e1e96e6"},
{file = "types_pyyaml-6.0.12.20241230.tar.gz", hash = "sha256:7f07622dbd34bb9c8b264fe860a17e0efcad00d50b5f27e93984909d9363498c"},
@@ -2494,6 +2707,7 @@ version = "4.12.2"
description = "Backported and Experimental Type Hints for Python 3.8+"
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "typing_extensions-4.12.2-py3-none-any.whl", hash = "sha256:04e5ca0351e0f3f85c6853954072df659d0d13fac324d0072316b67d7794700d"},
{file = "typing_extensions-4.12.2.tar.gz", hash = "sha256:1a7ead55c7e559dd4dee8856e3a88b41225abfe1ce8df57b7c13915fe121ffb8"},
@@ -2505,6 +2719,7 @@ version = "2.3.0"
description = "HTTP library with thread-safe connection pooling, file post, and more."
optional = false
python-versions = ">=3.9"
groups = ["main"]
files = [
{file = "urllib3-2.3.0-py3-none-any.whl", hash = "sha256:1cee9ad369867bfdbbb48b7dd50374c0967a0bb7710050facf0dd6911440e3df"},
{file = "urllib3-2.3.0.tar.gz", hash = "sha256:f8c5449b3cf0861679ce7e0503c7b44b5ec981bec0d1d3795a07f1ba96f0204d"},
@@ -2522,6 +2737,8 @@ version = "1.18.3"
description = "Yet another URL library"
optional = false
python-versions = ">=3.9"
groups = ["main"]
markers = "(python_version < \"3.10\" or python_version >= \"3.12\") and extra == \"graph\""
files = [
{file = "yarl-1.18.3-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:7df647e8edd71f000a5208fe6ff8c382a1de8edfbccdbbfe649d263de07d8c34"},
{file = "yarl-1.18.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:c69697d3adff5aa4f874b19c0e4ed65180ceed6318ec856ebc423aa5850d84f7"},
@@ -2618,6 +2835,8 @@ version = "0.23.0"
description = "Zstandard bindings for Python"
optional = false
python-versions = ">=3.8"
groups = ["main"]
markers = "extra == \"graph\""
files = [
{file = "zstandard-0.23.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:bf0a05b6059c0528477fba9054d09179beb63744355cab9f38059548fedd46a9"},
{file = "zstandard-0.23.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:fc9ca1c9718cb3b06634c7c8dec57d24e9438b2aa9a0f02b8bb36bf478538880"},
@@ -2728,6 +2947,6 @@ cffi = ["cffi (>=1.11)"]
graph = ["langchain-neo4j", "neo4j", "rank-bm25"]
[metadata]
lock-version = "2.0"
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "5848e23bdd7b453f938c9b5f6171866faa01bdcc2651bedb83ee9f4fe90e8bc8"
content-hash = "f7bee9b294566e32c580fef1940cbc6fcb7fd6ca2e695f23c719d134af4c3204"
+1 -1
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "mem0ai"
version = "0.1.69"
version = "0.1.74"
description = "Long-term memory for AI Agents"
authors = ["Mem0 <founders@mem0.ai>"]
exclude = [
+17
View File
@@ -0,0 +1,17 @@
from mem0.configs import prompts
def test_get_update_memory_messages():
retrieved_old_memory_dict = [{"id": "1", "text": "old memory 1"}]
response_content = ["new fact"]
custom_update_memory_prompt = "custom prompt determining memory update"
## When custom update memory prompt is provided
##
result = prompts.get_update_memory_messages(retrieved_old_memory_dict, response_content, custom_update_memory_prompt)
assert result.startswith(custom_update_memory_prompt)
## When custom update memory prompt is not provided
##
result = prompts.get_update_memory_messages(retrieved_old_memory_dict, response_content, None)
assert result.startswith(prompts.DEFAULT_UPDATE_MEMORY_PROMPT)
+53 -11
View File
@@ -20,10 +20,8 @@ def mock_openai_client():
yield mock_client
def test_generate_response(mock_openai_client):
config = BaseLlmConfig(
model=MODEL, temperature=TEMPERATURE, max_tokens=MAX_TOKENS, top_p=TOP_P
)
def test_generate_response_without_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."},
@@ -31,21 +29,67 @@ def test_generate_response(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 == "I'm doing well, thank you for asking!"
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."}
@pytest.mark.parametrize(
@@ -84,6 +128,4 @@ 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")
+64 -32
View File
@@ -16,47 +16,33 @@ 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(mock_deepseek_client):
config = BaseLlmConfig(
model="deepseek-chat", temperature=0.7, max_tokens=100, top_p=1.0
)
def test_generate_response_without_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."},
@@ -64,18 +50,64 @@ def test_generate_response(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."}
+79 -9
View File
@@ -17,9 +17,7 @@ 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."},
@@ -36,14 +34,86 @@ 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
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}}
),
)
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."}
+52 -8
View File
@@ -14,10 +14,8 @@ def mock_groq_client():
yield mock_client
def test_generate_response(mock_groq_client):
config = BaseLlmConfig(
model="llama3-70b-8192", temperature=0.7, max_tokens=100, top_p=1.0
)
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)
llm = GroqLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -25,18 +23,64 @@ def test_generate_response(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 == "I'm doing well, thank you for asking!"
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."}
+51 -11
View File
@@ -13,22 +13,17 @@ 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(mock_litellm):
def test_generate_response_without_tools(mock_litellm):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1)
llm = litellm.LiteLLM(config)
messages = [
@@ -37,9 +32,7 @@ def test_generate_response(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
@@ -49,3 +42,50 @@ def test_generate_response(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."}
+51 -16
View File
@@ -16,9 +16,7 @@ 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/"
@@ -26,9 +24,7 @@ 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 + "/"
@@ -36,19 +32,14 @@ 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(mock_openai_client):
def test_generate_response_without_tools(mock_openai_client):
config = BaseLlmConfig(model="gpt-4o", temperature=0.7, max_tokens=100, top_p=1.0)
llm = OpenAILLM(config)
messages = [
@@ -57,9 +48,7 @@ def test_generate_response(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)
@@ -68,3 +57,49 @@ def test_generate_response(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."}
+52 -11
View File
@@ -14,13 +14,8 @@ def mock_together_client():
yield mock_client
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,
)
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)
llm = TogetherLLM(config)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
@@ -28,18 +23,64 @@ def test_generate_response(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 == "I'm doing well, thank you for asking!"
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."}
+49
View File
@@ -30,6 +30,25 @@ def memory_instance():
config = MemoryConfig(version="v1.1")
config.graph_store.config = {"some_config": "value"}
return Memory(config)
@pytest.fixture
def memory_custom_instance():
with patch("mem0.utils.factory.EmbedderFactory") as mock_embedder, patch(
"mem0.utils.factory.VectorStoreFactory"
) as mock_vector_store, patch("mem0.utils.factory.LlmFactory") as mock_llm, patch(
"mem0.memory.telemetry.capture_event"
), patch("mem0.memory.graph_memory.MemoryGraph"):
mock_embedder.create.return_value = Mock()
mock_vector_store.create.return_value = Mock()
mock_llm.create.return_value = Mock()
config = MemoryConfig(
version="v1.1",
custom_fact_extraction_prompt="custom prompt extracting memory",
custom_update_memory_prompt="custom prompt determining memory update"
)
config.graph_store.config = {"some_config": "value"}
return Memory(config)
@pytest.mark.parametrize("version, enable_graph", [("v1.0", False), ("v1.1", True)])
@@ -239,3 +258,33 @@ def test_get_all(memory_instance, version, enable_graph, expected_result):
memory_instance.graph.get_all.assert_called_once_with({"user_id": "test_user"}, 100)
else:
memory_instance.graph.get_all.assert_not_called()
def test_custom_prompts(memory_custom_instance):
messages = [{"role": "user", "content": "Test message"}]
memory_custom_instance.llm.generate_response = Mock()
with patch("mem0.memory.main.parse_messages", return_value="Test message") as mock_parse_messages:
with patch("mem0.memory.main.get_update_memory_messages", return_value="custom update memory prompt") as mock_get_update_memory_messages:
memory_custom_instance.add(messages=messages, user_id="test_user")
## custom prompt
##
mock_parse_messages.assert_called_once_with(messages)
memory_custom_instance.llm.generate_response.assert_any_call(
messages=[
{"role": "system", "content": memory_custom_instance.config.custom_fact_extraction_prompt},
{"role": "user", "content": f"Input:\n{mock_parse_messages.return_value}"},
],
response_format={"type": "json_object"},
)
## custom update memory prompt
##
mock_get_update_memory_messages.assert_called_once_with([],[],memory_custom_instance.config.custom_update_memory_prompt)
memory_custom_instance.llm.generate_response.assert_any_call(
messages=[{"role": "user", "content": mock_get_update_memory_messages.return_value}],
response_format={"type": "json_object"},
)
+3 -4
View File
@@ -35,12 +35,11 @@ def test_search_vectors(chromadb_instance, mock_chromadb_client):
}
chromadb_instance.collection.query.return_value = mock_result
query = [[0.1, 0.2, 0.3]]
results = chromadb_instance.search(query=query, limit=2)
vectors = [[0.1, 0.2, 0.3]]
results = chromadb_instance.search(query="", vectors=vectors, limit=2)
chromadb_instance.collection.query.assert_called_once_with(query_embeddings=query, where=None, n_results=2)
chromadb_instance.collection.query.assert_called_once_with(query_embeddings=vectors, where=None, n_results=2)
print(results, type(results))
assert len(results) == 2
assert results[0].id == "id1"
assert results[0].score == 0.1
+21 -4
View File
@@ -1,6 +1,6 @@
import os
import unittest
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock, Mock, patch
import dotenv
@@ -196,8 +196,8 @@ class TestElasticsearchDB(unittest.TestCase):
self.client_mock.search.return_value = mock_response
# Perform search
query_vector = [0.1] * 1536
results = self.es_db.search(query=query_vector, limit=5)
vectors = [[0.1] * 1536]
results = self.es_db.search(query="", vectors=vectors, limit=5)
# Verify search call
self.client_mock.search.assert_called_once()
@@ -210,7 +210,7 @@ class TestElasticsearchDB(unittest.TestCase):
# Verify KNN query structure
self.assertIn("knn", body)
self.assertEqual(body["knn"]["field"], "vector")
self.assertEqual(body["knn"]["query_vector"], query_vector)
self.assertEqual(body["knn"]["query_vector"], vectors)
self.assertEqual(body["knn"]["k"], 5)
self.assertEqual(body["knn"]["num_candidates"], 10)
@@ -220,6 +220,23 @@ class TestElasticsearchDB(unittest.TestCase):
self.assertEqual(results[0].score, 0.8)
self.assertEqual(results[0].payload, {"key1": "value1"})
def test_custom_search_query(self):
# Mock custom search query
self.es_db.custom_search_query = Mock()
self.es_db.custom_search_query.return_value = {"custom_key": "custom_value"}
# Perform search
vectors = [[0.1] * 1536]
limit = 5
filters = {"key1": "value1"}
self.es_db.search(query="", vectors=vectors, limit=limit, filters=filters)
# Verify custom search query function was called
self.es_db.custom_search_query.assert_called_once_with(vectors, limit, filters)
# Verify custom search query was used
self.client_mock.search.assert_called_once_with(index=self.es_db.collection_name, body={"custom_key": "custom_value"})
def test_get(self):
# Mock get response with correct structure
mock_response = {
+30 -4
View File
@@ -5,7 +5,7 @@ from unittest.mock import MagicMock, patch
import dotenv
try:
from opensearchpy import OpenSearch
from opensearchpy import OpenSearch, AWSV4SignerAuth
except ImportError:
raise ImportError(
"OpenSearch requires extra dependencies. Install with `pip install opensearch-py`"
@@ -126,15 +126,15 @@ class TestOpenSearchDB(unittest.TestCase):
def test_search(self):
mock_response = {"hits": {"hits": [{"_id": "id1", "_score": 0.8, "_source": {"vector": [0.1] * 1536, "metadata": {"key1": "value1"}}}]}}
self.client_mock.search.return_value = mock_response
query_vector = [0.1] * 1536
results = self.os_db.search(query=query_vector, limit=5)
vectors = [[0.1] * 1536]
results = self.os_db.search(query="", vectors=vectors, limit=5)
self.client_mock.search.assert_called_once()
search_args = self.client_mock.search.call_args[1]
self.assertEqual(search_args["index"], "test_collection")
body = search_args["body"]
self.assertIn("knn", body["query"])
self.assertIn("vector", body["query"]["knn"])
self.assertEqual(body["query"]["knn"]["vector"]["vector"], query_vector)
self.assertEqual(body["query"]["knn"]["vector"]["vector"], vectors)
self.assertEqual(body["query"]["knn"]["vector"]["k"], 5)
self.assertEqual(len(results), 1)
self.assertEqual(results[0].id, "id1")
@@ -148,3 +148,29 @@ class TestOpenSearchDB(unittest.TestCase):
def test_delete_col(self):
self.os_db.delete_col()
self.client_mock.indices.delete.assert_called_once_with(index="test_collection")
def test_init_with_http_auth(self):
mock_credentials = MagicMock()
mock_signer = AWSV4SignerAuth(mock_credentials, "us-east-1", "es")
with patch('mem0.vector_stores.opensearch.OpenSearch') as mock_opensearch:
test_db = OpenSearchDB(
host="localhost",
port=9200,
collection_name="test_collection",
embedding_model_dims=1536,
http_auth=mock_signer,
verify_certs=True,
use_ssl=True,
auto_create_index=False
)
# Verify OpenSearch was initialized with correct params
mock_opensearch.assert_called_once_with(
hosts=[{"host": "localhost", "port": 9200}],
http_auth=mock_signer,
use_ssl=True,
verify_certs=True,
connection_class=unittest.mock.ANY
)
+120
View File
@@ -0,0 +1,120 @@
from unittest.mock import MagicMock
import pytest
from mem0.vector_stores.pinecone import PineconeDB
@pytest.fixture
def mock_pinecone_client():
client = MagicMock()
client.Index.return_value = MagicMock()
client.list_indexes.return_value.names.return_value = []
return client
@pytest.fixture
def pinecone_db(mock_pinecone_client):
return PineconeDB(
collection_name="test_index",
embedding_model_dims=128,
client=mock_pinecone_client,
api_key="fake_api_key",
environment="us-west1-gcp",
serverless_config=None,
pod_config=None,
hybrid_search=False,
metric="cosine",
batch_size=100,
extra_params=None
)
def test_create_col_existing_index(mock_pinecone_client):
# Set up the mock before creating the PineconeDB object
mock_pinecone_client.list_indexes.return_value.names.return_value = ["test_index"]
pinecone_db = PineconeDB(
collection_name="test_index",
embedding_model_dims=128,
client=mock_pinecone_client,
api_key="fake_api_key",
environment="us-west1-gcp",
serverless_config=None,
pod_config=None,
hybrid_search=False,
metric="cosine",
batch_size=100,
extra_params=None
)
# Reset the mock to verify it wasn't called during the test
mock_pinecone_client.create_index.reset_mock()
pinecone_db.create_col(128, "cosine")
mock_pinecone_client.create_index.assert_not_called()
def test_create_col_new_index(pinecone_db, mock_pinecone_client):
mock_pinecone_client.list_indexes.return_value.names.return_value = []
pinecone_db.create_col(128, "cosine")
mock_pinecone_client.create_index.assert_called()
def test_insert_vectors(pinecone_db):
vectors = [[0.1] * 128, [0.2] * 128]
payloads = [{"name": "vector1"}, {"name": "vector2"}]
ids = ["id1", "id2"]
pinecone_db.insert(vectors, payloads, ids)
pinecone_db.index.upsert.assert_called()
def test_search_vectors(pinecone_db):
pinecone_db.index.query.return_value.matches = [{"id": "id1", "score": 0.9, "metadata": {"name": "vector1"}}]
results = pinecone_db.search([0.1] * 128, limit=1)
assert len(results) == 1
assert results[0].id == "id1"
assert results[0].score == 0.9
def test_update_vector(pinecone_db):
pinecone_db.update("id1", vector=[0.5] * 128, payload={"name": "updated"})
pinecone_db.index.upsert.assert_called()
def test_get_vector_found(pinecone_db):
# Looking at the _parse_output method, it expects a Vector object
# or a list of dictionaries, not a dictionary with an 'id' field
# Create a mock Vector object
from pinecone.data.dataclasses.vector import Vector
mock_vector = Vector(
id="id1",
values=[0.1] * 128,
metadata={"name": "vector1"}
)
# Mock the fetch method to return the mock response object
mock_response = MagicMock()
mock_response.vectors = {"id1": mock_vector}
pinecone_db.index.fetch.return_value = mock_response
result = pinecone_db.get("id1")
assert result is not None
assert result.id == "id1"
assert result.payload == {"name": "vector1"}
def test_delete_vector(pinecone_db):
pinecone_db.delete("id1")
pinecone_db.index.delete.assert_called_with(ids=["id1"])
def test_get_vector_not_found(pinecone_db):
pinecone_db.index.fetch.return_value.vectors = {}
result = pinecone_db.get("id1")
assert result is None
def test_list_cols(pinecone_db):
pinecone_db.list_cols()
pinecone_db.client.list_indexes.assert_called()
def test_delete_col(pinecone_db):
pinecone_db.delete_col()
pinecone_db.client.delete_index.assert_called_with("test_index")
def test_col_info(pinecone_db):
pinecone_db.col_info()
pinecone_db.client.describe_index.assert_called_with("test_index")
+3 -3
View File
@@ -50,15 +50,15 @@ class TestQdrant(unittest.TestCase):
self.assertEqual(points[0].payload, payloads[0])
def test_search(self):
query_vector = [0.1, 0.2]
vectors = [[0.1, 0.2]]
mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"})
self.client_mock.query_points.return_value = MagicMock(points=[mock_point])
results = self.qdrant.search(query=query_vector, limit=1)
results = self.qdrant.search(query="", vectors=vectors, limit=1)
self.client_mock.query_points.assert_called_once_with(
collection_name="test_collection",
query=query_vector,
query=vectors,
query_filter=None,
limit=1,
)
+3 -3
View File
@@ -77,12 +77,12 @@ def test_search_vectors(supabase_instance, mock_collection):
]
mock_collection.query.return_value = mock_results
query = [0.1, 0.2, 0.3]
vectors = [[0.1, 0.2, 0.3]]
filters = {"category": "test"}
results = supabase_instance.search(query=query, limit=2, filters=filters)
results = supabase_instance.search(query="", vectors=vectors, limit=2, filters=filters)
mock_collection.query.assert_called_once_with(
data=query,
data=vectors,
limit=2,
filters={"category": {"$eq": "test"}},
include_metadata=True,
@@ -73,12 +73,12 @@ def test_insert_vectors(vector_store, mock_vertex_ai):
def test_search_vectors(vector_store, mock_vertex_ai):
"""Test searching vectors with filters"""
query = [0.1, 0.2, 0.3]
vectors = [[0.1, 0.2, 0.3]]
filters = {"user_id": "test_user"}
mock_datapoint = Mock()
mock_datapoint.datapoint_id = "test-id"
mock_datapoint.feature_vector = query
mock_datapoint.feature_vector = vectors
mock_restrict = Mock()
mock_restrict.namespace = "user_id"
@@ -96,11 +96,11 @@ def test_search_vectors(vector_store, mock_vertex_ai):
mock_vertex_ai['endpoint'].find_neighbors.return_value = [[mock_neighbor]]
results = vector_store.search(query=query, filters=filters, limit=1)
results = vector_store.search(query="", vectors=vectors, filters=filters, limit=1)
mock_vertex_ai['endpoint'].find_neighbors.assert_called_once_with(
deployed_index_id=vector_store.deployment_index_id,
queries=[query],
queries=[vectors],
num_neighbors=1,
filter=[Namespace("user_id", ["test_user"], [])],
return_full_datapoint=True
+2 -2
View File
@@ -147,8 +147,8 @@ class TestWeaviateDB(unittest.TestCase):
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)
vectors = [[0.1] * 1536]
results = self.weaviate_db.search(query="", vectors=vectors, limit=5)
mock_hybrid.assert_called_once()