diff --git a/docs/components/vectordbs/dbs/turbopuffer.mdx b/docs/components/vectordbs/dbs/turbopuffer.mdx
new file mode 100644
index 000000000..893c4b6f0
--- /dev/null
+++ b/docs/components/vectordbs/dbs/turbopuffer.mdx
@@ -0,0 +1,75 @@
+[Turbopuffer](https://turbopuffer.com) is a serverless vector database optimized for low-latency search at scale. It offers cost-effective vector storage with native metadata filtering.
+
+### Usage
+
+```python
+import os
+from mem0 import Memory
+
+os.environ["OPENAI_API_KEY"] = "sk-xx"
+os.environ["TURBOPUFFER_API_KEY"] = "tpuf_xxxxxxxxxxxx"
+
+config = {
+ "vector_store": {
+ "provider": "turbopuffer",
+ "config": {
+ "collection_name": "movie_preferences",
+ "embedding_model_dims": 1536,
+ "region": "gcp-us-central1",
+ }
+ }
+}
+
+m = Memory.from_config(config)
+
+messages = [
+ {"role": "user", "content": "I'm planning to watch a movie tonight. Any recommendations?"},
+ {"role": "assistant", "content": "How about thriller movies? They can be quite engaging."},
+ {"role": "user", "content": "I'm not a big fan of thrillers but I love sci-fi."},
+ {"role": "assistant", "content": "Got it! I'll suggest sci-fi movies instead."}
+]
+
+m.add(messages, user_id="alice", metadata={"category": "movies"})
+
+# Search memories
+results = m.search(query="sci-fi recommendations", user_id="alice")
+```
+
+### Config
+
+Here are the parameters available for configuring Turbopuffer:
+
+| Parameter | Description | Default Value |
+| --- | --- | --- |
+| `collection_name` | Name of the namespace/collection | `mem0` |
+| `embedding_model_dims` | Dimensions of the embedding model (must match your chosen embedding model) | `1536` |
+| `api_key` | Turbopuffer API key | Environment variable: `TURBOPUFFER_API_KEY` |
+| `region` | Turbopuffer region | `gcp-us-central1` |
+| `distance_metric` | Distance metric for vector similarity (`cosine_distance` or `euclidean_squared`) | `cosine_distance` |
+| `batch_size` | Batch size for bulk operations | `100` |
+| `extra_params` | Additional parameters for the Turbopuffer client | `None` |
+
+### Regions
+
+| Region | Location |
+| --- | --- |
+| `gcp-us-central1` | Iowa, USA (Default) |
+| `aws-us-west-2` | Oregon, USA |
+
+### Config Example
+
+```python
+config = {
+ "vector_store": {
+ "provider": "turbopuffer",
+ "config": {
+ "collection_name": "my_memories",
+ "embedding_model_dims": 1536,
+ "api_key": "tpuf_xxxxxxxxxxxx",
+ "region": "aws-us-west-2",
+ "distance_metric": "cosine_distance",
+ "batch_size": 200,
+ }
+ }
+}
+```
diff --git a/docs/components/vectordbs/overview.mdx b/docs/components/vectordbs/overview.mdx
index f7d3a9e8a..dbc06e7ed 100644
--- a/docs/components/vectordbs/overview.mdx
+++ b/docs/components/vectordbs/overview.mdx
@@ -33,6 +33,7 @@ See the list of supported vector databases below.
+
## Usage
diff --git a/docs/docs.json b/docs/docs.json
index 623d42fbb..c09af9b83 100644
--- a/docs/docs.json
+++ b/docs/docs.json
@@ -231,7 +231,8 @@
"components/vectordbs/dbs/cassandra",
"components/vectordbs/dbs/s3_vectors",
"components/vectordbs/dbs/databricks",
- "components/vectordbs/dbs/neptune_analytics"
+ "components/vectordbs/dbs/neptune_analytics",
+ "components/vectordbs/dbs/turbopuffer"
]
}
]
diff --git a/mem0/configs/vector_stores/turbopuffer.py b/mem0/configs/vector_stores/turbopuffer.py
new file mode 100644
index 000000000..1d43715ca
--- /dev/null
+++ b/mem0/configs/vector_stores/turbopuffer.py
@@ -0,0 +1,45 @@
+import os
+from typing import Any, Dict, Optional
+
+from pydantic import BaseModel, ConfigDict, Field, model_validator
+
+
+class TurbopufferConfig(BaseModel):
+ collection_name: str = Field("mem0", description="Name of the namespace/collection")
+ embedding_model_dims: int = Field(1536, description="Dimensions of the embedding model")
+ api_key: Optional[str] = Field(None, description="API key for Turbopuffer")
+ region: str = Field("gcp-us-central1", description="Turbopuffer region (e.g., 'gcp-us-central1', 'aws-us-west-2')")
+ distance_metric: str = Field(
+ "cosine_distance",
+ description="Distance metric for vector similarity ('cosine_distance' or 'euclidean_squared')",
+ )
+ batch_size: int = Field(100, description="Batch size for bulk operations")
+ extra_params: Optional[Dict[str, Any]] = Field(
+ None,
+ description="Additional parameters for Turbopuffer client",
+ )
+
+ @model_validator(mode="before")
+ @classmethod
+ def check_api_key(cls, values: Dict[str, Any]) -> Dict[str, Any]:
+ api_key = values.get("api_key")
+ if not api_key and "TURBOPUFFER_API_KEY" not in os.environ:
+ raise ValueError(
+ "Either 'api_key' must be provided or TURBOPUFFER_API_KEY environment variable must be set."
+ )
+ 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)}. "
+ f"Please input only the following fields: {', '.join(allowed_fields)}"
+ )
+ return values
+
+ model_config = ConfigDict(arbitrary_types_allowed=True)
diff --git a/mem0/utils/factory.py b/mem0/utils/factory.py
index 87a76c885..85d8755b9 100644
--- a/mem0/utils/factory.py
+++ b/mem0/utils/factory.py
@@ -188,6 +188,7 @@ class VectorStoreFactory:
"baidu": "mem0.vector_stores.baidu.BaiduDB",
"cassandra": "mem0.vector_stores.cassandra.CassandraDB",
"neptune": "mem0.vector_stores.neptune_analytics.NeptuneAnalyticsVector",
+ "turbopuffer": "mem0.vector_stores.turbopuffer.TurbopufferDB",
}
@classmethod
diff --git a/mem0/vector_stores/configs.py b/mem0/vector_stores/configs.py
index d08bae37a..32459dddc 100644
--- a/mem0/vector_stores/configs.py
+++ b/mem0/vector_stores/configs.py
@@ -34,6 +34,7 @@ class VectorStoreConfig(BaseModel):
"faiss": "FAISSConfig",
"langchain": "LangchainConfig",
"s3_vectors": "S3VectorsConfig",
+ "turbopuffer": "TurbopufferConfig",
}
@model_validator(mode="after")
diff --git a/mem0/vector_stores/turbopuffer.py b/mem0/vector_stores/turbopuffer.py
new file mode 100644
index 000000000..a6708746f
--- /dev/null
+++ b/mem0/vector_stores/turbopuffer.py
@@ -0,0 +1,337 @@
+import logging
+import os
+from typing import Any, Dict, List, Optional, Union
+
+try:
+ from turbopuffer import Turbopuffer as TurbopufferClient
+except ImportError:
+ raise ImportError(
+ "Turbopuffer requires extra dependencies. Install with `pip install turbopuffer`"
+ ) from None
+
+from pydantic import BaseModel
+
+from mem0.vector_stores.base import VectorStoreBase
+
+logger = logging.getLogger(__name__)
+
+
+class OutputData(BaseModel):
+ id: Optional[str]
+ score: Optional[float]
+ payload: Optional[Dict]
+
+
+class TurbopufferDB(VectorStoreBase):
+ def __init__(
+ self,
+ collection_name: str,
+ embedding_model_dims: int,
+ api_key: Optional[str] = None,
+ region: str = "gcp-us-central1",
+ distance_metric: str = "cosine_distance",
+ batch_size: int = 100,
+ extra_params: Optional[Dict[str, Any]] = None,
+ ):
+ """
+ Initialize the Turbopuffer vector store.
+
+ Args:
+ collection_name (str): Name of the namespace/collection.
+ embedding_model_dims (int): Dimensions of the embedding model.
+ api_key (str, optional): API key for Turbopuffer. Defaults to None.
+ region (str, optional): Turbopuffer region. Defaults to "gcp-us-central1".
+ distance_metric (str, optional): Distance metric for vector similarity.
+ Options: "cosine_distance" or "euclidean_squared". Defaults to "cosine_distance".
+ batch_size (int, optional): Batch size for operations. Defaults to 100.
+ extra_params (Dict, optional): Additional parameters for Turbopuffer client. Defaults to None.
+ """
+ api_key = api_key or os.environ.get("TURBOPUFFER_API_KEY")
+ if not api_key:
+ raise ValueError(
+ "Turbopuffer API key must be provided either as a parameter or via TURBOPUFFER_API_KEY environment variable"
+ )
+
+ params = extra_params or {}
+ params["region"] = region
+
+ self.client = TurbopufferClient(api_key=api_key, **params)
+ self.collection_name = collection_name
+ self.embedding_model_dims = embedding_model_dims
+ self.distance_metric = distance_metric
+ self.batch_size = batch_size
+
+ self.namespace = self.client.namespace(self.collection_name)
+
+ def create_col(self, name=None, vector_size=None, distance=None):
+ """
+ Create a new namespace in Turbopuffer.
+ Namespaces are created implicitly on first upsert, so this is a no-op.
+ """
+ pass
+
+ def insert(
+ self,
+ vectors: List[List[float]],
+ payloads: Optional[List[Dict]] = None,
+ ids: Optional[List[Union[str, int]]] = None,
+ ):
+ """
+ Insert vectors into the namespace.
+
+ 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 namespace {self.collection_name}")
+
+ if ids is None:
+ ids = [str(i) for i in range(len(vectors))]
+
+ for i in range(0, len(vectors), self.batch_size):
+ batch_end = i + self.batch_size
+ rows = []
+ for j in range(i, min(batch_end, len(vectors))):
+ row = {}
+ if payloads and payloads[j]:
+ row.update(payloads[j])
+ row["id"] = str(ids[j])
+ row["vector"] = vectors[j]
+ rows.append(row)
+
+ self.namespace.write(
+ upsert_rows=rows,
+ distance_metric=self.distance_metric,
+ )
+
+ def _parse_output(self, rows) -> List[OutputData]:
+ """
+ Parse the output data from Turbopuffer query results.
+
+ Args:
+ rows: List of Row objects from Turbopuffer query.
+
+ Returns:
+ List[OutputData]: Parsed output data.
+ """
+ results = []
+ for row in rows:
+ row_dict = row.model_dump()
+ row_id = str(row_dict.pop("id"))
+ dist = row_dict.pop("$dist", None)
+ row_dict.pop("vector", None)
+
+ score = 1 - dist if dist is not None else None
+
+ results.append(OutputData(
+ id=row_id,
+ score=score,
+ payload=row_dict,
+ ))
+ return results
+
+ def _convert_filters(self, filters: Optional[Dict]):
+ """
+ Convert mem0 filters to Turbopuffer filter format.
+
+ Turbopuffer filters use tuple format: ("And", (("field", "Op", value), ...))
+ """
+ if not filters:
+ return None
+
+ conditions = []
+ for key, value in filters.items():
+ if isinstance(value, dict):
+ if "gte" in value:
+ conditions.append((key, "Gte", value["gte"]))
+ if "lte" in value:
+ conditions.append((key, "Lte", value["lte"]))
+ else:
+ conditions.append((key, "Eq", value))
+
+ if not conditions:
+ return None
+ if len(conditions) == 1:
+ return conditions[0]
+ return ("And", tuple(conditions))
+
+ 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 text (unused in vector search, kept for interface consistency).
+ vectors (list): Query vector to search with.
+ 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.
+ """
+ query_params = {
+ "rank_by": ("vector", "ANN", vectors),
+ "top_k": limit,
+ "include_attributes": True,
+ }
+
+ tpuf_filters = self._convert_filters(filters)
+ if tpuf_filters is not None:
+ query_params["filters"] = tpuf_filters
+
+ response = self.namespace.query(**query_params)
+ return self._parse_output(response.rows or [])
+
+ 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.namespace.write(deletes=[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.
+ """
+ if vector is not None:
+ row = {}
+ if payload:
+ row.update(payload)
+ row["id"] = str(vector_id)
+ row["vector"] = vector
+ self.namespace.write(
+ upsert_rows=[row],
+ distance_metric=self.distance_metric,
+ )
+ elif payload is not None:
+ row = dict(payload)
+ row["id"] = str(vector_id)
+ self.namespace.write(patch_rows=[row])
+
+ def get(self, vector_id: Union[str, int]) -> Optional[OutputData]:
+ """
+ Retrieve a vector by ID.
+
+ Args:
+ vector_id (Union[str, int]): ID of the vector to retrieve.
+
+ Returns:
+ OutputData: Retrieved vector data, or None if not found.
+ """
+ try:
+ response = self.namespace.query(
+ top_k=1,
+ rank_by=("vector", "ANN", [0.0] * self.embedding_model_dims),
+ filters=("id", "Eq", str(vector_id)),
+ include_attributes=True,
+ )
+ rows = response.rows or []
+ if rows:
+ return self._parse_output(rows)[0]
+ return None
+ except Exception as e:
+ logger.error(f"Error retrieving vector {vector_id}: {e}")
+ return None
+
+ def list_cols(self) -> list:
+ """
+ List all namespaces.
+
+ Returns:
+ list: List of namespace summaries.
+ """
+ try:
+ result = []
+ for ns in self.client.namespaces():
+ result.append(ns)
+ return result
+ except Exception as e:
+ logger.error(f"Error listing namespaces: {e}")
+ return []
+
+ def delete_col(self):
+ """Delete the entire namespace."""
+ try:
+ self.namespace.delete_all()
+ logger.info(f"Namespace {self.collection_name} deleted successfully")
+ except Exception as e:
+ logger.error(f"Error deleting namespace {self.collection_name}: {e}")
+
+ def col_info(self) -> Dict:
+ """
+ Get information about the namespace.
+
+ Returns:
+ dict: Namespace metadata.
+ """
+ try:
+ metadata = self.namespace.metadata()
+ return {
+ "name": self.collection_name,
+ "approx_row_count": metadata.approx_row_count,
+ "approx_logical_bytes": metadata.approx_logical_bytes,
+ "created_at": str(metadata.created_at),
+ "updated_at": str(metadata.updated_at),
+ }
+ except Exception:
+ return {"name": self.collection_name}
+
+ def list(self, filters: Optional[Dict] = None, limit: int = 100) -> list:
+ """
+ List vectors in the namespace with optional filtering.
+
+ Args:
+ filters (dict, optional): Filters to apply. Defaults to None.
+ limit (int, optional): Number of vectors to return. Defaults to 100.
+
+ Returns:
+ list: Wrapped list of OutputData objects ([[results]]).
+ """
+ query_params = {
+ "rank_by": ("vector", "ANN", [0.0] * self.embedding_model_dims),
+ "top_k": limit,
+ "include_attributes": True,
+ }
+
+ tpuf_filters = self._convert_filters(filters)
+ if tpuf_filters is not None:
+ query_params["filters"] = tpuf_filters
+
+ try:
+ response = self.namespace.query(**query_params)
+ results = self._parse_output(response.rows or [])
+ except Exception as e:
+ logger.error(f"Error listing vectors: {e}")
+ results = []
+ return [results]
+
+ def count(self) -> int:
+ """
+ Get approximate count of vectors in the namespace.
+
+ Returns:
+ int: Approximate number of vectors.
+ """
+ try:
+ metadata = self.namespace.metadata()
+ return metadata.approx_row_count
+ except Exception:
+ return 0
+
+ def reset(self):
+ """Reset the namespace by deleting all vectors."""
+ self.delete_col()
diff --git a/tests/vector_stores/test_turbopuffer.py b/tests/vector_stores/test_turbopuffer.py
new file mode 100644
index 000000000..7e677ed60
--- /dev/null
+++ b/tests/vector_stores/test_turbopuffer.py
@@ -0,0 +1,725 @@
+from datetime import datetime
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+pytest.importorskip("turbopuffer", reason="turbopuffer not installed")
+
+from turbopuffer.types import Row
+
+from mem0.vector_stores.turbopuffer import OutputData, TurbopufferDB
+
+
+def _make_row(id, dist=None, vector=None, **attributes):
+ """Helper to create a turbopuffer Row with extra attributes."""
+ data = {"id": id, "vector": vector}
+ if dist is not None:
+ data["$dist"] = dist
+ data.update(attributes)
+ return Row.model_validate(data)
+
+
+@pytest.fixture
+def mock_client():
+ with patch("mem0.vector_stores.turbopuffer.TurbopufferClient") as MockClient:
+ client_instance = MagicMock()
+ namespace_instance = MagicMock()
+ client_instance.namespace.return_value = namespace_instance
+ MockClient.return_value = client_instance
+ yield client_instance, namespace_instance
+
+
+@pytest.fixture
+def db(mock_client):
+ _, _ = mock_client
+ return TurbopufferDB(
+ collection_name="test_ns",
+ embedding_model_dims=4,
+ api_key="tpuf_test_key",
+ region="gcp-us-central1",
+ distance_metric="cosine_distance",
+ batch_size=2,
+ )
+
+
+# ── Initialization ──────────────────────────────────────────────────
+
+
+class TestInit:
+ def test_init_with_api_key(self, mock_client):
+ client_instance, namespace_instance = mock_client
+ db = TurbopufferDB(
+ collection_name="my_ns",
+ embedding_model_dims=128,
+ api_key="tpuf_key",
+ region="gcp-us-central1",
+ )
+ assert db.collection_name == "my_ns"
+ assert db.embedding_model_dims == 128
+ assert db.distance_metric == "cosine_distance"
+ assert db.batch_size == 100
+ client_instance.namespace.assert_called_with("my_ns")
+
+ def test_init_with_env_var(self, mock_client):
+ with patch.dict("os.environ", {"TURBOPUFFER_API_KEY": "tpuf_env_key"}):
+ db = TurbopufferDB(
+ collection_name="test",
+ embedding_model_dims=4,
+ region="gcp-us-central1",
+ )
+ assert db.collection_name == "test"
+
+ def test_init_without_api_key_raises(self, mock_client):
+ with patch.dict("os.environ", {}, clear=True):
+ # Remove TURBOPUFFER_API_KEY if it exists
+ import os
+ os.environ.pop("TURBOPUFFER_API_KEY", None)
+ with pytest.raises(ValueError, match="API key must be provided"):
+ TurbopufferDB(
+ collection_name="test",
+ embedding_model_dims=4,
+ region="gcp-us-central1",
+ )
+
+ def test_init_with_extra_params(self, mock_client):
+ client_instance, _ = mock_client
+ with patch("mem0.vector_stores.turbopuffer.TurbopufferClient") as MockClient:
+ MockClient.return_value = client_instance
+ TurbopufferDB(
+ collection_name="test",
+ embedding_model_dims=4,
+ api_key="key",
+ region="aws-us-west-2",
+ extra_params={"compression": True},
+ )
+ MockClient.assert_called_with(
+ api_key="key",
+ region="aws-us-west-2",
+ compression=True,
+ )
+
+ def test_init_default_region(self, mock_client):
+ client_instance, _ = mock_client
+ with patch("mem0.vector_stores.turbopuffer.TurbopufferClient") as MockClient:
+ MockClient.return_value = client_instance
+ TurbopufferDB(
+ collection_name="test",
+ embedding_model_dims=4,
+ api_key="key",
+ )
+ MockClient.assert_called_with(
+ api_key="key",
+ region="gcp-us-central1",
+ )
+
+
+# ── create_col ───────────────────────────────────────────────────────
+
+
+class TestCreateCol:
+ def test_create_col_is_noop(self, db):
+ # Should not raise or call anything
+ result = db.create_col()
+ assert result is None
+
+ def test_create_col_with_args_is_noop(self, db):
+ result = db.create_col(name="x", vector_size=128, distance="cosine")
+ assert result is None
+
+
+# ── insert ───────────────────────────────────────────────────────────
+
+
+class TestInsert:
+ def test_insert_with_ids_and_payloads(self, db):
+ vectors = [[0.1, 0.2, 0.3, 0.4], [0.5, 0.6, 0.7, 0.8]]
+ payloads = [{"data": "hello", "user_id": "u1"}, {"data": "world", "user_id": "u2"}]
+ ids = ["id1", "id2"]
+
+ db.insert(vectors, payloads, ids)
+
+ db.namespace.write.assert_called_once()
+ call_kwargs = db.namespace.write.call_args[1]
+ rows = call_kwargs["upsert_rows"]
+ assert len(rows) == 2
+ assert rows[0]["id"] == "id1"
+ assert rows[0]["vector"] == [0.1, 0.2, 0.3, 0.4]
+ assert rows[0]["data"] == "hello"
+ assert rows[0]["user_id"] == "u1"
+ assert rows[1]["id"] == "id2"
+ assert call_kwargs["distance_metric"] == "cosine_distance"
+
+ def test_insert_without_ids_generates_ids(self, db):
+ vectors = [[0.1, 0.2, 0.3, 0.4]]
+ db.insert(vectors)
+
+ db.namespace.write.assert_called_once()
+ call_kwargs = db.namespace.write.call_args[1]
+ assert call_kwargs["upsert_rows"][0]["id"] == "0"
+
+ def test_insert_without_payloads(self, db):
+ vectors = [[0.1, 0.2, 0.3, 0.4]]
+ ids = ["id1"]
+ db.insert(vectors, ids=ids)
+
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["upsert_rows"][0]
+ assert row["id"] == "id1"
+ assert row["vector"] == [0.1, 0.2, 0.3, 0.4]
+ assert len(row) == 2 # only id and vector
+
+ def test_insert_payload_does_not_overwrite_id_or_vector(self, db):
+ """Payload with 'id' or 'vector' keys must not overwrite the actual values."""
+ vectors = [[0.1, 0.2, 0.3, 0.4]]
+ payloads = [{"id": "fake_id", "vector": "fake_vector", "data": "hello"}]
+ ids = ["real_id"]
+ db.insert(vectors, payloads, ids)
+
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["upsert_rows"][0]
+ assert row["id"] == "real_id"
+ assert row["vector"] == [0.1, 0.2, 0.3, 0.4]
+ assert row["data"] == "hello"
+
+ def test_insert_batching(self, db):
+ """batch_size=2, so 3 vectors should produce 2 write calls."""
+ vectors = [[0.1] * 4, [0.2] * 4, [0.3] * 4]
+ ids = ["a", "b", "c"]
+ db.insert(vectors, ids=ids)
+
+ assert db.namespace.write.call_count == 2
+ first_call = db.namespace.write.call_args_list[0][1]
+ second_call = db.namespace.write.call_args_list[1][1]
+ assert len(first_call["upsert_rows"]) == 2
+ assert len(second_call["upsert_rows"]) == 1
+
+
+# ── _parse_output ────────────────────────────────────────────────────
+
+
+class TestParseOutput:
+ def test_parse_rows_with_dist_and_attributes(self, db):
+ rows = [
+ _make_row("id1", dist=0.1, data="hello", user_id="u1"),
+ _make_row("id2", dist=0.3, data="world", user_id="u2"),
+ ]
+ results = db._parse_output(rows)
+
+ assert len(results) == 2
+ assert results[0].id == "id1"
+ assert results[0].score == pytest.approx(0.9)
+ assert results[0].payload == {"data": "hello", "user_id": "u1"}
+ assert results[1].id == "id2"
+ assert results[1].score == pytest.approx(0.7)
+ assert results[1].payload == {"data": "world", "user_id": "u2"}
+
+ def test_parse_rows_without_dist(self, db):
+ rows = [_make_row("id1", data="hello")]
+ results = db._parse_output(rows)
+
+ assert results[0].score is None
+ assert results[0].payload == {"data": "hello"}
+
+ def test_parse_rows_strips_vector_and_id(self, db):
+ rows = [_make_row("id1", dist=0.0, vector=[1.0, 2.0, 3.0, 4.0], data="test")]
+ results = db._parse_output(rows)
+
+ assert "vector" not in results[0].payload
+ assert "id" not in results[0].payload
+ assert "$dist" not in results[0].payload
+ assert results[0].id == "id1"
+
+ def test_parse_empty_rows(self, db):
+ assert db._parse_output([]) == []
+
+
+# ── _convert_filters ─────────────────────────────────────────────────
+
+
+class TestConvertFilters:
+ def test_none_filters(self, db):
+ assert db._convert_filters(None) is None
+
+ def test_empty_filters(self, db):
+ assert db._convert_filters({}) is None
+
+ def test_single_eq_filter(self, db):
+ result = db._convert_filters({"user_id": "u1"})
+ assert result == ("user_id", "Eq", "u1")
+
+ def test_multiple_eq_filters(self, db):
+ result = db._convert_filters({"user_id": "u1", "agent_id": "a1"})
+ assert result[0] == "And"
+ conditions = result[1]
+ assert ("user_id", "Eq", "u1") in conditions
+ assert ("agent_id", "Eq", "a1") in conditions
+
+ def test_range_filter_gte_lte(self, db):
+ result = db._convert_filters({"score": {"gte": 0.5, "lte": 1.0}})
+ assert result[0] == "And"
+ conditions = result[1]
+ assert ("score", "Gte", 0.5) in conditions
+ assert ("score", "Lte", 1.0) in conditions
+
+ def test_range_filter_gte_only(self, db):
+ result = db._convert_filters({"score": {"gte": 0.5}})
+ assert result == ("score", "Gte", 0.5)
+
+ def test_range_filter_lte_only(self, db):
+ result = db._convert_filters({"score": {"lte": 1.0}})
+ assert result == ("score", "Lte", 1.0)
+
+ def test_mixed_eq_and_range_filters(self, db):
+ result = db._convert_filters({"user_id": "u1", "score": {"gte": 0.5}})
+ assert result[0] == "And"
+ conditions = result[1]
+ assert ("user_id", "Eq", "u1") in conditions
+ assert ("score", "Gte", 0.5) in conditions
+
+
+# ── search ───────────────────────────────────────────────────────────
+
+
+class TestSearch:
+ def test_search_basic(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = [
+ _make_row("id1", dist=0.1, data="hello"),
+ _make_row("id2", dist=0.2, data="world"),
+ ]
+ db.namespace.query.return_value = mock_response
+
+ results = db.search("test query", [0.1, 0.2, 0.3, 0.4], limit=2)
+
+ db.namespace.query.assert_called_once_with(
+ rank_by=("vector", "ANN", [0.1, 0.2, 0.3, 0.4]),
+ top_k=2,
+ include_attributes=True,
+ )
+ assert len(results) == 2
+ assert results[0].id == "id1"
+ assert results[0].score == pytest.approx(0.9)
+
+ def test_search_with_filters(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = [_make_row("id1", dist=0.1, data="hello")]
+ db.namespace.query.return_value = mock_response
+
+ results = db.search(
+ "query", [0.1, 0.2, 0.3, 0.4], limit=5, filters={"user_id": "u1"}
+ )
+
+ call_kwargs = db.namespace.query.call_args[1]
+ assert call_kwargs["filters"] == ("user_id", "Eq", "u1")
+ assert len(results) == 1
+
+ def test_search_no_results(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = None
+ db.namespace.query.return_value = mock_response
+
+ results = db.search("query", [0.1, 0.2, 0.3, 0.4])
+ assert results == []
+
+ def test_search_empty_rows(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = []
+ db.namespace.query.return_value = mock_response
+
+ results = db.search("query", [0.1, 0.2, 0.3, 0.4])
+ assert results == []
+
+
+# ── delete ───────────────────────────────────────────────────────────
+
+
+class TestDelete:
+ def test_delete_by_string_id(self, db):
+ db.delete("id1")
+ db.namespace.write.assert_called_once_with(deletes=["id1"])
+
+ def test_delete_by_int_id(self, db):
+ db.delete(123)
+ db.namespace.write.assert_called_once_with(deletes=["123"])
+
+
+# ── update ───────────────────────────────────────────────────────────
+
+
+class TestUpdate:
+ def test_update_with_vector_and_payload(self, db):
+ db.update("id1", vector=[0.5, 0.6, 0.7, 0.8], payload={"data": "updated"})
+
+ db.namespace.write.assert_called_once()
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["upsert_rows"][0]
+ assert row["id"] == "id1"
+ assert row["vector"] == [0.5, 0.6, 0.7, 0.8]
+ assert row["data"] == "updated"
+ assert call_kwargs["distance_metric"] == "cosine_distance"
+
+ def test_update_vector_only(self, db):
+ db.update("id1", vector=[0.5, 0.6, 0.7, 0.8])
+
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["upsert_rows"][0]
+ assert row["id"] == "id1"
+ assert row["vector"] == [0.5, 0.6, 0.7, 0.8]
+ assert "data" not in row
+
+ def test_update_payload_only_uses_patch(self, db):
+ """Payload-only updates should use patch_rows, not upsert_rows."""
+ db.update("id1", vector=None, payload={"data": "patched", "user_id": "u1"})
+
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["patch_rows"][0]
+ assert row["id"] == "id1"
+ assert row["data"] == "patched"
+ assert row["user_id"] == "u1"
+ assert "upsert_rows" not in call_kwargs
+
+ def test_update_payload_does_not_overwrite_id(self, db):
+ """Payload with an 'id' key must not overwrite the actual vector ID."""
+ db.update("real_id", vector=None, payload={"id": "fake_id", "data": "test"})
+ call_kwargs = db.namespace.write.call_args[1]
+ row = call_kwargs["patch_rows"][0]
+ assert row["id"] == "real_id"
+
+ def test_update_nothing(self, db):
+ """Neither vector nor payload: should not call write."""
+ db.update("id1")
+ db.namespace.write.assert_not_called()
+
+
+# ── get ──────────────────────────────────────────────────────────────
+
+
+class TestGet:
+ def test_get_found(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = [_make_row("id1", dist=0.0, data="hello", user_id="u1")]
+ db.namespace.query.return_value = mock_response
+
+ result = db.get("id1")
+
+ db.namespace.query.assert_called_once_with(
+ top_k=1,
+ rank_by=("vector", "ANN", [0.0, 0.0, 0.0, 0.0]),
+ filters=("id", "Eq", "id1"),
+ include_attributes=True,
+ )
+ assert result is not None
+ assert result.id == "id1"
+ assert result.payload["data"] == "hello"
+ assert result.payload["user_id"] == "u1"
+
+ def test_get_not_found(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = []
+ db.namespace.query.return_value = mock_response
+
+ result = db.get("nonexistent")
+ assert result is None
+
+ def test_get_none_rows(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = None
+ db.namespace.query.return_value = mock_response
+
+ result = db.get("id1")
+ assert result is None
+
+ def test_get_handles_exception(self, db):
+ db.namespace.query.side_effect = Exception("API error")
+ result = db.get("id1")
+ assert result is None
+
+
+# ── list_cols ────────────────────────────────────────────────────────
+
+
+class TestListCols:
+ def test_list_cols(self, db):
+ ns1 = MagicMock()
+ ns1.id = "ns1"
+ ns2 = MagicMock()
+ ns2.id = "ns2"
+ db.client.namespaces.return_value = [ns1, ns2]
+
+ result = db.list_cols()
+ assert len(result) == 2
+ db.client.namespaces.assert_called_once()
+
+
+# ── delete_col ───────────────────────────────────────────────────────
+
+
+class TestDeleteCol:
+ def test_delete_col(self, db):
+ db.delete_col()
+ db.namespace.delete_all.assert_called_once()
+
+ def test_delete_col_handles_error(self, db):
+ db.namespace.delete_all.side_effect = Exception("API error")
+ # Should not raise
+ db.delete_col()
+
+
+# ── col_info ─────────────────────────────────────────────────────────
+
+
+class TestColInfo:
+ def test_col_info(self, db):
+ mock_metadata = MagicMock()
+ mock_metadata.approx_row_count = 42
+ mock_metadata.approx_logical_bytes = 1024
+ mock_metadata.created_at = datetime(2025, 1, 1)
+ mock_metadata.updated_at = datetime(2025, 6, 1)
+ db.namespace.metadata.return_value = mock_metadata
+
+ info = db.col_info()
+ assert info["name"] == "test_ns"
+ assert info["approx_row_count"] == 42
+ assert info["approx_logical_bytes"] == 1024
+
+ def test_col_info_handles_error(self, db):
+ db.namespace.metadata.side_effect = Exception("not found")
+ info = db.col_info()
+ assert info == {"name": "test_ns"}
+
+
+# ── list ─────────────────────────────────────────────────────────────
+
+
+class TestList:
+ def test_list_returns_wrapped_format(self, db):
+ """list() must return [[results]] for compatibility with main.py."""
+ mock_response = MagicMock()
+ mock_response.rows = [
+ _make_row("id1", dist=0.1, data="hello"),
+ _make_row("id2", dist=0.2, data="world"),
+ ]
+ db.namespace.query.return_value = mock_response
+
+ result = db.list()
+
+ # Must be wrapped: result[0] is the actual list
+ assert isinstance(result, list)
+ assert isinstance(result[0], list)
+ assert len(result[0]) == 2
+ assert result[0][0].id == "id1"
+ assert result[0][1].id == "id2"
+
+ def test_list_with_filters(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = [_make_row("id1", dist=0.1, data="hello")]
+ db.namespace.query.return_value = mock_response
+
+ db.list(filters={"user_id": "u1"}, limit=50)
+
+ call_kwargs = db.namespace.query.call_args[1]
+ assert call_kwargs["filters"] == ("user_id", "Eq", "u1")
+ assert call_kwargs["top_k"] == 50
+
+ def test_list_empty(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = None
+ db.namespace.query.return_value = mock_response
+
+ result = db.list()
+ assert result == [[]]
+
+ def test_list_uses_zero_vector(self, db):
+ mock_response = MagicMock()
+ mock_response.rows = []
+ db.namespace.query.return_value = mock_response
+
+ db.list()
+ call_kwargs = db.namespace.query.call_args[1]
+ assert call_kwargs["rank_by"] == ("vector", "ANN", [0.0, 0.0, 0.0, 0.0])
+
+ def test_list_compatible_with_main_py_get_all(self, db):
+ """Simulate how main.py _get_all_from_vector_store unwraps list()."""
+ mock_response = MagicMock()
+ mock_response.rows = [_make_row("id1", dist=0.1, data="hello")]
+ db.namespace.query.return_value = mock_response
+
+ memories_result = db.list(filters={"user_id": "u1"})
+
+ # Reproduce main.py unwrapping logic
+ first_element = memories_result[0]
+ if isinstance(first_element, (list, tuple)):
+ actual_memories = first_element
+ else:
+ actual_memories = memories_result
+
+ assert isinstance(actual_memories, list)
+ assert len(actual_memories) == 1
+ assert actual_memories[0].id == "id1"
+ assert actual_memories[0].payload["data"] == "hello"
+
+ def test_list_compatible_with_main_py_delete_all(self, db):
+ """Simulate how main.py delete_all uses list()[0]."""
+ mock_response = MagicMock()
+ mock_response.rows = [
+ _make_row("id1", dist=0.1, data="hello"),
+ _make_row("id2", dist=0.2, data="world"),
+ ]
+ db.namespace.query.return_value = mock_response
+
+ result = db.list(filters={"user_id": "u1"})
+ memories = result[0]
+
+ assert isinstance(memories, list)
+ for mem in memories:
+ assert hasattr(mem, "id")
+ assert hasattr(mem, "payload")
+
+
+# ── count ────────────────────────────────────────────────────────────
+
+
+class TestCount:
+ def test_count(self, db):
+ mock_metadata = MagicMock()
+ mock_metadata.approx_row_count = 100
+ db.namespace.metadata.return_value = mock_metadata
+
+ assert db.count() == 100
+
+ def test_count_handles_error(self, db):
+ db.namespace.metadata.side_effect = Exception("error")
+ assert db.count() == 0
+
+
+# ── reset ────────────────────────────────────────────────────────────
+
+
+class TestReset:
+ def test_reset_calls_delete_all(self, db):
+ db.reset()
+ db.namespace.delete_all.assert_called_once()
+
+
+# ── Config ───────────────────────────────────────────────────────────
+
+
+class TestConfig:
+ def test_config_valid(self):
+ with patch.dict("os.environ", {"TURBOPUFFER_API_KEY": "key"}):
+ from mem0.configs.vector_stores.turbopuffer import TurbopufferConfig
+
+ config = TurbopufferConfig()
+ assert config.collection_name == "mem0"
+ assert config.embedding_model_dims == 1536
+ assert config.distance_metric == "cosine_distance"
+ assert config.batch_size == 100
+ assert config.region == "gcp-us-central1"
+
+ def test_config_custom_values(self):
+ from mem0.configs.vector_stores.turbopuffer import TurbopufferConfig
+
+ config = TurbopufferConfig(
+ collection_name="custom",
+ embedding_model_dims=768,
+ api_key="tpuf_key",
+ region="aws-us-west-2",
+ distance_metric="euclidean_squared",
+ batch_size=50,
+ )
+ assert config.collection_name == "custom"
+ assert config.embedding_model_dims == 768
+ assert config.api_key == "tpuf_key"
+ assert config.region == "aws-us-west-2"
+ assert config.distance_metric == "euclidean_squared"
+ assert config.batch_size == 50
+
+ def test_config_rejects_extra_fields(self):
+ from mem0.configs.vector_stores.turbopuffer import TurbopufferConfig
+
+ with pytest.raises(ValueError, match="Extra fields not allowed"):
+ TurbopufferConfig(api_key="key", unknown_field="value")
+
+ def test_config_requires_api_key(self):
+ from mem0.configs.vector_stores.turbopuffer import TurbopufferConfig
+
+ with patch.dict("os.environ", {}, clear=True):
+ import os
+ os.environ.pop("TURBOPUFFER_API_KEY", None)
+ with pytest.raises(ValueError, match="api_key"):
+ TurbopufferConfig()
+
+ def test_config_accepts_env_var(self):
+ from mem0.configs.vector_stores.turbopuffer import TurbopufferConfig
+
+ with patch.dict("os.environ", {"TURBOPUFFER_API_KEY": "env_key"}):
+ config = TurbopufferConfig()
+ assert config.api_key is None # not set explicitly, but env var is present
+
+
+# ── Factory Registration ─────────────────────────────────────────────
+
+
+class TestFactoryRegistration:
+ def test_vector_store_factory_has_turbopuffer(self):
+ from mem0.utils.factory import VectorStoreFactory
+
+ assert "turbopuffer" in VectorStoreFactory.provider_to_class
+ assert VectorStoreFactory.provider_to_class["turbopuffer"] == "mem0.vector_stores.turbopuffer.TurbopufferDB"
+
+ def test_vector_store_config_has_turbopuffer(self):
+ from mem0.vector_stores.configs import VectorStoreConfig
+
+ config = VectorStoreConfig.__private_attributes__["_provider_configs"].default
+ assert "turbopuffer" in config
+ assert config["turbopuffer"] == "TurbopufferConfig"
+
+ def test_config_validation_pipeline(self):
+ """Test that VectorStoreConfig correctly resolves turbopuffer config."""
+ from mem0.vector_stores.configs import VectorStoreConfig
+
+ with patch.dict("os.environ", {"TURBOPUFFER_API_KEY": "key"}):
+ config = VectorStoreConfig(
+ provider="turbopuffer",
+ config={"collection_name": "test", "region": "gcp-us-central1"},
+ )
+ assert config.config.collection_name == "test"
+ assert config.config.region == "gcp-us-central1"
+ assert config.config.embedding_model_dims == 1536
+
+
+# ── OutputData ───────────────────────────────────────────────────────
+
+
+class TestOutputData:
+ def test_output_data_has_required_fields(self):
+ od = OutputData(id="test", score=0.9, payload={"data": "hello"})
+ assert od.id == "test"
+ assert od.score == 0.9
+ assert od.payload == {"data": "hello"}
+
+ def test_output_data_nullable_fields(self):
+ od = OutputData(id=None, score=None, payload=None)
+ assert od.id is None
+ assert od.score is None
+ assert od.payload is None
+
+ def test_output_data_payload_access_pattern(self):
+ """Test the exact access pattern used by main.py."""
+ od = OutputData(
+ id="mem1",
+ score=0.95,
+ payload={
+ "data": "User likes sci-fi movies",
+ "hash": "abc123",
+ "created_at": "2025-01-01",
+ "updated_at": "2025-06-01",
+ "user_id": "alice",
+ "agent_id": "agent1",
+ },
+ )
+ assert od.payload.get("data", "") == "User likes sci-fi movies"
+ assert od.payload.get("hash") == "abc123"
+ assert od.payload.get("created_at") == "2025-01-01"
+ assert od.payload.get("user_id") == "alice"
+ assert od.payload.get("run_id") is None # not present, should return None