Refactor reranking and output format for the OSS and platform
This commit is contained in:
@@ -19,14 +19,13 @@ client.get_all(version="v1", output_format="v1.1")
|
||||
|
||||
**After (1.0.0):**
|
||||
```python
|
||||
# v1.1 API is default (v1.0 is deprecated)
|
||||
# v1.1 format is default (v1.0 is deprecated)
|
||||
memory = Memory() # Defaults to v1.1 format
|
||||
|
||||
# Client API with correct versioning behavior:
|
||||
client.add(messages) # Uses v1 API endpoint
|
||||
client.add(messages, output_format="v1.1") # Uses v1.1 API endpoint
|
||||
client.search(query) # Uses v2 API endpoint
|
||||
client.get_all() # Uses v2 API endpoint
|
||||
client.add(messages) # Uses v1 API endpoint, returns v1.1 format
|
||||
client.search(query) # Uses v2 API endpoint, returns v1.1 format
|
||||
client.get_all() # Uses v2 API endpoint, returns v1.1 format
|
||||
```
|
||||
|
||||
### 2. API Versioning Strategy Clarification
|
||||
@@ -74,7 +73,7 @@ memories = memory.get_all()
|
||||
|
||||
# v1.0 format still works but shows deprecation warning
|
||||
memory_v1 = Memory(config=MemoryConfig(version="v1.0"))
|
||||
result = memory_v1.add(messages) # Returns raw list (with warning)
|
||||
result = memory_v1.add(messages) # Returns raw list [{...}] (with warning)
|
||||
```
|
||||
|
||||
## Migration Steps
|
||||
|
||||
+30
-10
@@ -184,7 +184,7 @@ class MemoryClient:
|
||||
return response.json()
|
||||
|
||||
@api_error_handler
|
||||
def get_all(self, **kwargs) -> List[Dict[str, Any]]:
|
||||
def get_all(self, **kwargs) -> Dict[str, Any]:
|
||||
"""Retrieve all memories, with optional filtering.
|
||||
|
||||
Args:
|
||||
@@ -192,7 +192,7 @@ class MemoryClient:
|
||||
app_id, top_k, page, page_size, version).
|
||||
|
||||
Returns:
|
||||
A list of dictionaries containing memories.
|
||||
A dictionary containing memories in v1.1 format: {"results": [...]}
|
||||
|
||||
Raises:
|
||||
APIError: If the API request fails.
|
||||
@@ -223,10 +223,15 @@ class MemoryClient:
|
||||
"sync_type": "sync",
|
||||
},
|
||||
)
|
||||
return response.json()
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
|
||||
@api_error_handler
|
||||
def search(self, query: str, **kwargs) -> List[Dict[str, Any]]:
|
||||
def search(self, query: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Search memories based on a query.
|
||||
|
||||
Args:
|
||||
@@ -235,7 +240,7 @@ class MemoryClient:
|
||||
top_k, filters, version.
|
||||
|
||||
Returns:
|
||||
A list of dictionaries containing search results.
|
||||
A dictionary containing search results in v1.1 format: {"results": [...]}
|
||||
|
||||
Raises:
|
||||
APIError: If the API request fails.
|
||||
@@ -262,7 +267,12 @@ class MemoryClient:
|
||||
"sync_type": "sync",
|
||||
},
|
||||
)
|
||||
return response.json()
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
|
||||
@api_error_handler
|
||||
def update(
|
||||
@@ -1018,7 +1028,7 @@ class AsyncMemoryClient:
|
||||
return response.json()
|
||||
|
||||
@api_error_handler
|
||||
async def get_all(self, **kwargs) -> List[Dict[str, Any]]:
|
||||
async def get_all(self, **kwargs) -> Dict[str, Any]:
|
||||
params = self._prepare_params(kwargs)
|
||||
# Handle version parameter for get operations (default to v2)
|
||||
version = kwargs.pop("version", "v2")
|
||||
@@ -1045,10 +1055,15 @@ class AsyncMemoryClient:
|
||||
"sync_type": "async",
|
||||
},
|
||||
)
|
||||
return response.json()
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
|
||||
@api_error_handler
|
||||
async def search(self, query: str, **kwargs) -> List[Dict[str, Any]]:
|
||||
async def search(self, query: str, **kwargs) -> Dict[str, Any]:
|
||||
payload = {"query": query}
|
||||
params = self._prepare_params(kwargs)
|
||||
# Handle version parameter for search operations (default to v2)
|
||||
@@ -1071,7 +1086,12 @@ class AsyncMemoryClient:
|
||||
"sync_type": "async",
|
||||
},
|
||||
)
|
||||
return response.json()
|
||||
result = response.json()
|
||||
|
||||
# Ensure v1.1 format (wrap raw list if needed)
|
||||
if isinstance(result, list):
|
||||
return {"results": result}
|
||||
return result
|
||||
|
||||
@api_error_handler
|
||||
async def update(
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
from .base import BaseRerankerConfig
|
||||
from .cohere import CohereRerankerConfig
|
||||
from .sentence_transformer import SentenceTransformerRerankerConfig
|
||||
from .config import RerankerConfig
|
||||
|
||||
__all__ = ["BaseRerankerConfig", "CohereRerankerConfig", "SentenceTransformerRerankerConfig", "RerankerConfig"]
|
||||
@@ -1,8 +1,8 @@
|
||||
from abc import ABC
|
||||
from typing import Optional
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class BaseRerankerConfig(ABC):
|
||||
class BaseRerankerConfig(BaseModel):
|
||||
"""
|
||||
Base configuration for rerankers with only common parameters.
|
||||
Provider-specific configurations should be handled by separate config classes.
|
||||
@@ -11,20 +11,7 @@ class BaseRerankerConfig(ABC):
|
||||
For provider-specific parameters, use the appropriate provider config class.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: Optional[str] = None,
|
||||
api_key: Optional[str] = None,
|
||||
top_k: Optional[int] = None,
|
||||
):
|
||||
"""
|
||||
Initialize a base configuration class instance for the reranker.
|
||||
|
||||
Args:
|
||||
model (str, optional): The reranker model to use.
|
||||
api_key (str, optional): The API key for the reranker service.
|
||||
top_k (int, optional): Maximum number of documents to return after reranking.
|
||||
"""
|
||||
self.model = model
|
||||
self.api_key = api_key
|
||||
self.top_k = top_k
|
||||
provider: str = Field(description="The reranker provider to use")
|
||||
model: Optional[str] = Field(default=None, description="The reranker model to use")
|
||||
api_key: Optional[str] = Field(default=None, description="The API key for the reranker service")
|
||||
top_k: Optional[int] = Field(default=None, description="Maximum number of documents to return after reranking")
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Optional
|
||||
from pydantic import Field
|
||||
|
||||
from mem0.configs.rerankers.base import BaseRerankerConfig
|
||||
|
||||
@@ -9,26 +10,7 @@ class CohereRerankerConfig(BaseRerankerConfig):
|
||||
Inherits from BaseRerankerConfig and adds Cohere-specific settings.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
# Base parameters
|
||||
model: Optional[str] = "rerank-english-v3.0",
|
||||
api_key: Optional[str] = None,
|
||||
top_k: Optional[int] = None,
|
||||
# Cohere-specific parameters
|
||||
return_documents: bool = False,
|
||||
max_chunks_per_doc: Optional[int] = None,
|
||||
):
|
||||
"""
|
||||
Initialize Cohere reranker configuration.
|
||||
|
||||
Args:
|
||||
model (str): The Cohere rerank model to use.
|
||||
api_key (str, optional): The Cohere API key.
|
||||
top_k (int, optional): Maximum number of documents to return after reranking.
|
||||
return_documents (bool): Whether to return the document texts in the response.
|
||||
max_chunks_per_doc (int, optional): Maximum number of chunks per document.
|
||||
"""
|
||||
super().__init__(model=model, api_key=api_key, top_k=top_k)
|
||||
self.return_documents = return_documents
|
||||
self.max_chunks_per_doc = max_chunks_per_doc
|
||||
provider: str = Field(default="cohere", description="The reranker provider")
|
||||
model: Optional[str] = Field(default="rerank-english-v3.0", description="The Cohere rerank model to use")
|
||||
return_documents: bool = Field(default=False, description="Whether to return the document texts in the response")
|
||||
max_chunks_per_doc: Optional[int] = Field(default=None, description="Maximum number of chunks per document")
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Optional
|
||||
from pydantic import Field
|
||||
|
||||
from mem0.configs.rerankers.base import BaseRerankerConfig
|
||||
|
||||
@@ -9,29 +10,8 @@ class SentenceTransformerRerankerConfig(BaseRerankerConfig):
|
||||
Inherits from BaseRerankerConfig and adds Sentence Transformer-specific settings.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
# Base parameters
|
||||
model: Optional[str] = "cross-encoder/ms-marco-MiniLM-L-6-v2",
|
||||
api_key: Optional[str] = None, # Not used for sentence transformers
|
||||
top_k: Optional[int] = None,
|
||||
# Sentence Transformer-specific parameters
|
||||
device: Optional[str] = None,
|
||||
batch_size: int = 32,
|
||||
show_progress_bar: bool = False,
|
||||
):
|
||||
"""
|
||||
Initialize Sentence Transformer reranker configuration.
|
||||
|
||||
Args:
|
||||
model (str): The cross-encoder model name to use.
|
||||
api_key (str, optional): Not used for sentence transformers.
|
||||
top_k (int, optional): Maximum number of documents to return after reranking.
|
||||
device (str, optional): Device to run the model on ('cpu', 'cuda', etc.).
|
||||
batch_size (int): Batch size for processing documents.
|
||||
show_progress_bar (bool): Whether to show progress bar during processing.
|
||||
"""
|
||||
super().__init__(model=model, api_key=api_key, top_k=top_k)
|
||||
self.device = device
|
||||
self.batch_size = batch_size
|
||||
self.show_progress_bar = show_progress_bar
|
||||
provider: str = Field(default="sentence_transformer", description="The reranker provider")
|
||||
model: Optional[str] = Field(default="cross-encoder/ms-marco-MiniLM-L-6-v2", description="The cross-encoder model name to use")
|
||||
device: Optional[str] = Field(default=None, description="Device to run the model on ('cpu', 'cuda', etc.)")
|
||||
batch_size: int = Field(default=32, description="Batch size for processing documents")
|
||||
show_progress_bar: bool = Field(default=False, description="Whether to show progress bar during processing")
|
||||
Reference in New Issue
Block a user