Refactor reranking and output format for the OSS and platform

This commit is contained in:
parshvadaftari
2025-09-02 02:20:24 +05:30
parent 61e8668584
commit 5182c8f311
6 changed files with 52 additions and 90 deletions
+5 -6
View File
@@ -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
View File
@@ -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(
-6
View File
@@ -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"]
+6 -19
View File
@@ -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")
+5 -23
View File
@@ -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")
+6 -26
View File
@@ -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")