354 lines
11 KiB
Python
354 lines
11 KiB
Python
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
|
|
|
|
# Maps mem0 filter operators to their Turbopuffer equivalents.
|
|
OPERATOR_MAP = {
|
|
"eq": "Eq",
|
|
"ne": "NotEq",
|
|
"gt": "Gt",
|
|
"gte": "Gte",
|
|
"lt": "Lt",
|
|
"lte": "Lte",
|
|
"in": "In",
|
|
"nin": "NotIn",
|
|
}
|
|
|
|
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):
|
|
for op, operand in value.items():
|
|
tpuf_op = self.OPERATOR_MAP.get(op)
|
|
if tpuf_op is None:
|
|
raise ValueError(
|
|
f"Unsupported filter operator '{op}' for field '{key}'. "
|
|
f"Supported operators: {sorted(self.OPERATOR_MAP)}"
|
|
)
|
|
conditions.append((key, tpuf_op, operand))
|
|
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], top_k: 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.
|
|
top_k (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": top_k,
|
|
"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, top_k: int = 100) -> list:
|
|
"""
|
|
List vectors in the namespace with optional filtering.
|
|
|
|
Args:
|
|
filters (dict, optional): Filters to apply. Defaults to None.
|
|
top_k (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": top_k,
|
|
"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()
|