Compare commits
26 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ff30cb8ddd | |||
| 2e853c3d22 | |||
| 3cc7013fde | |||
| 8e6a08aa83 | |||
| 8b9a8e5825 | |||
| afc630272d | |||
| e33008e3a4 | |||
| 7b516328a8 | |||
| 6d5889d98f | |||
| ee66e0c954 | |||
| 1aed611539 | |||
| 6c2b131d6e | |||
| 2ffe9922f3 | |||
| 540ada489b | |||
| 65cffa0369 | |||
| 9f937943ba | |||
| 51a68bf7c5 | |||
| 92541d8955 | |||
| 0e0be18ecc | |||
| 66d3f9b93c | |||
| b8f40f728f | |||
| 00a2ea9ff0 | |||
| 9545836469 | |||
| 3acd9e20da | |||
| d48ecd52ef | |||
| 2fbea7705b |
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
@@ -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
@@ -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"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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" />
|
||||
+11
-11
@@ -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 |
|
||||
@@ -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>
|
||||
|
||||
@@ -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" />
|
||||
|
||||
@@ -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"
|
||||
```
|
||||
|
||||
|
||||
@@ -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 |
@@ -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"
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
}
|
||||
},
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -23,6 +23,8 @@ export type {
|
||||
Message,
|
||||
AllUsers,
|
||||
User,
|
||||
FeedbackPayload,
|
||||
Feedback,
|
||||
} from "./mem0.types";
|
||||
|
||||
// Export telemetry types
|
||||
|
||||
@@ -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 };
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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);
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -13,8 +13,9 @@ export interface Message {
|
||||
}
|
||||
|
||||
export interface EmbeddingConfig {
|
||||
apiKey: string;
|
||||
apiKey?: string;
|
||||
model?: string;
|
||||
url?: string;
|
||||
}
|
||||
|
||||
export interface VectorStoreConfig {
|
||||
|
||||
@@ -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}`);
|
||||
}
|
||||
|
||||
@@ -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];
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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.
|
||||
"""
|
||||
"""
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"qdrant": "QdrantConfig",
|
||||
"chroma": "ChromaDbConfig",
|
||||
"pgvector": "PGVectorConfig",
|
||||
"pinecone": "PineconeConfig",
|
||||
"milvus": "MilvusDBConfig",
|
||||
"azure_ai_search": "AzureAISearchConfig",
|
||||
"redis": "RedisDBConfig",
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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=["*"],
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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 = [
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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."}
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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."}
|
||||
|
||||
@@ -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"},
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
)
|
||||
@@ -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")
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user