from __future__ import annotations
import json
import uuid
from typing import Any, Dict, List, Optional, cast
from llama_index.core.schema import BaseNode, TextNode
from llama_index.core.vector_stores.types import (
BasePydanticVectorStore,
VectorStoreQuery,
VectorStoreQueryResult,
)
class AgentDBVectorStore(BasePydanticVectorStore):
stores_text: bool = True
flat_metadata: bool = False
_db_path: str
_collection_name: str
_dimension: int
_db: Any
_collection: Any
def __init__(
self,
db_path: str,
collection_name: str,
dimension: int,
**kwargs: Any,
) -> None:
try:
import agentdb as _agentdb
except ImportError as exc: raise ImportError(
"The 'datacules-agentdb' package is required. "
"Install it with: pip install datacules-agentdb"
) from exc
super().__init__(**kwargs)
object.__setattr__(self, "_db_path", db_path)
object.__setattr__(self, "_collection_name", collection_name)
object.__setattr__(self, "_dimension", dimension)
db = _agentdb.AgentDB.open(db_path)
object.__setattr__(self, "_db", db)
object.__setattr__(self, "_collection", db.collection(collection_name, dimension))
@classmethod
def class_name(cls) -> str:
return "AgentDBVectorStore"
@property
def client(self) -> Any:
return self._db
def add(self, nodes: List[BaseNode], **kwargs: Any) -> List[str]:
if not nodes:
return []
entries = []
ids: List[str] = []
for node in nodes:
node_id = node.node_id or str(uuid.uuid4())
ids.append(node_id)
embedding = node.get_embedding()
if embedding is None:
raise ValueError(
f"Node {node_id!r} has no embedding. "
"Embed the nodes before calling add()."
)
stored_meta: Dict[str, Any] = {}
if isinstance(node, TextNode):
stored_meta["_text"] = node.text
else:
try:
stored_meta["_text"] = node.get_content()
except Exception:
stored_meta["_text"] = ""
stored_meta["_node_type"] = type(node).__name__
stored_meta["_metadata"] = node.metadata or {}
if node.ref_doc_id:
stored_meta["_ref_doc_id"] = node.ref_doc_id
entries.append(
{
"id": node_id,
"vector": [float(v) for v in embedding],
"metadata": stored_meta,
}
)
self._collection.upsert_batch(entries)
return ids
def delete(self, ref_doc_id: str, **kwargs: Any) -> None:
node_id: Optional[str] = kwargs.get("node_id")
if node_id is not None:
self._collection.delete(node_id)
return
zero_vec = [0.0] * self._dimension
try:
results = self._collection.search(zero_vec, top_k=10_000)
except Exception:
results = []
for result in results:
meta = result.metadata or {}
if meta.get("_ref_doc_id") == ref_doc_id:
self._collection.delete(result.id)
def query(self, query: VectorStoreQuery, **kwargs: Any) -> VectorStoreQueryResult:
if query.query_embedding is None:
raise ValueError(
"query.query_embedding must be set before calling query(). "
"Use an embedder to generate the query vector first."
)
top_k = query.similarity_top_k or 4
raw_results = self._collection.search(
[float(v) for v in query.query_embedding],
top_k=top_k,
)
result_nodes: List[BaseNode] = []
similarities: List[float] = []
ids: List[str] = []
for result in raw_results:
meta = result.metadata or {}
text = meta.get("_text", "")
node_metadata: Dict[str, Any] = meta.get("_metadata", {})
node = TextNode(
id_=result.id,
text=text,
metadata=node_metadata,
)
result_nodes.append(node)
similarities.append(float(result.score))
ids.append(result.id)
return VectorStoreQueryResult(
nodes=result_nodes,
similarities=similarities,
ids=ids,
)
def __len__(self) -> int:
return self._collection.count()
def __repr__(self) -> str:
return (
f"AgentDBVectorStore("
f"db_path={self._db_path!r}, "
f"collection={self._collection_name!r}, "
f"dim={self._dimension})"
)