feat: integrate turbopuffer as vector database provider (#4428)
Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com> Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
@@ -33,6 +33,7 @@ See the list of supported vector databases below.
|
||||
<Card title="LangChain" href="/components/vectordbs/dbs/langchain"></Card>
|
||||
<Card title="Amazon S3 Vectors" href="/components/vectordbs/dbs/s3_vectors"></Card>
|
||||
<Card title="Databricks" href="/components/vectordbs/dbs/databricks"></Card>
|
||||
<Card title="Turbopuffer" href="/components/vectordbs/dbs/turbopuffer"></Card>
|
||||
</CardGroup>
|
||||
|
||||
## Usage
|
||||
|
||||
+2
-1
@@ -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"
|
||||
]
|
||||
}
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -34,6 +34,7 @@ class VectorStoreConfig(BaseModel):
|
||||
"faiss": "FAISSConfig",
|
||||
"langchain": "LangchainConfig",
|
||||
"s3_vectors": "S3VectorsConfig",
|
||||
"turbopuffer": "TurbopufferConfig",
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user