diff --git a/MIGRATION_GUIDE_v1.0.md b/MIGRATION_GUIDE_v1.0.md index db921169f..920d2e900 100644 --- a/MIGRATION_GUIDE_v1.0.md +++ b/MIGRATION_GUIDE_v1.0.md @@ -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 diff --git a/mem0/client/main.py b/mem0/client/main.py index c2545420e..c6b8f9240 100644 --- a/mem0/client/main.py +++ b/mem0/client/main.py @@ -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( diff --git a/mem0/configs/rerankers/__init__.py b/mem0/configs/rerankers/__init__.py index 8544b7c5b..e69de29bb 100644 --- a/mem0/configs/rerankers/__init__.py +++ b/mem0/configs/rerankers/__init__.py @@ -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"] \ No newline at end of file diff --git a/mem0/configs/rerankers/base.py b/mem0/configs/rerankers/base.py index cc2be05a6..c48a0a1a0 100644 --- a/mem0/configs/rerankers/base.py +++ b/mem0/configs/rerankers/base.py @@ -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 \ No newline at end of file + 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") \ No newline at end of file diff --git a/mem0/configs/rerankers/cohere.py b/mem0/configs/rerankers/cohere.py index 52230aedc..d05dc15dc 100644 --- a/mem0/configs/rerankers/cohere.py +++ b/mem0/configs/rerankers/cohere.py @@ -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 \ No newline at end of file + 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") \ No newline at end of file diff --git a/mem0/configs/rerankers/sentence_transformer.py b/mem0/configs/rerankers/sentence_transformer.py index 6e320be94..6339cb74d 100644 --- a/mem0/configs/rerankers/sentence_transformer.py +++ b/mem0/configs/rerankers/sentence_transformer.py @@ -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 \ No newline at end of file + 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") \ No newline at end of file