sane-assistant / src /memory /vector_db.py
agenthinkmesh's picture
Upload project files
c13f797 verified
Raw History Blame Contribute Delete
4.82 kB
import chromadb
import google.generativeai as genai
import os
import json
import sys
from typing import List, Dict, Any
class VectorDB:
def __init__(self, path: str = "./.sane/memory/chroma_db", collection_name: str = "sane_memory"):
self.client = chromadb.PersistentClient(path=path)
self.collection = self.client.get_or_create_collection(name=collection_name)
self.model = "gemini-embedding-001"
genai.configure(api_key=os.getenv("GEMINI_API_KEY"))
def _generate_embeddings(self, texts: List[str]) -> List[List[float]]:
if not texts:
return []
try:
result = genai.embed_content(
model=self.model,
content=texts,
task_type="SEMANTIC_SIMILARITY"
)
return [embedding.values for embedding in result["embeddings"]]
except Exception as e:
print(f"Error generating embeddings: {e}", file=sys.stderr)
return []
def add_documents(self, documents: List[str], metadatas: List[Dict[str, Any]], ids: List[str]):
embeddings = self._generate_embeddings(documents)
if not embeddings:
print("No embeddings generated, skipping adding documents.", file=sys.stderr)
return
filtered_documents = []
filtered_metadatas = []
filtered_ids = []
filtered_embeddings = []
for i, emb in enumerate(embeddings):
if emb:
filtered_documents.append(documents[i])
filtered_metadatas.append(metadatas[i])
filtered_ids.append(ids[i])
filtered_embeddings.append(emb)
if not filtered_documents:
print("All embedding generations failed, no documents added.", file=sys.stderr)
return
self.collection.add(
embeddings=filtered_embeddings,
documents=filtered_documents,
metadatas=filtered_metadatas,
ids=filtered_ids
)
print(f"Added {len(filtered_documents)} documents to ChromaDB.", file=sys.stderr)
def query(self, query_text: str, n_results: int = 5) -> Dict[str, Any]:
query_embedding = self._generate_embeddings([query_text])
if not query_embedding or not query_embedding[0]:
print("Could not generate embedding for query.", file=sys.stderr)
return {"documents": [], "metadatas": [], "distances": [], "ids": []}
results = self.collection.query(
query_embeddings=query_embedding[0],
n_results=n_results,
include=["documents", "metadatas", "distances", "ids"]
)
return results
def get_all_documents(self) -> Dict[str, Any]:
results = self.collection.get(
ids=self.collection.get()["ids"],
include=["documents", "metadatas", "ids"]
)
return results
def delete_documents(self, ids: List[str]):
self.collection.delete(ids=ids)
print(f"Deleted {len(ids)} documents from ChromaDB.", file=sys.stderr)
def count(self) -> int:
return self.collection.count()
if __name__ == "__main__":
db = VectorDB()
command = sys.argv[1]
if command == "add":
doc = sys.argv[2]
metadata_str = sys.argv[3]
doc_id = sys.argv[4]
metadata = json.loads(metadata_str)
db.add_documents(documents=[doc], metadatas=[metadata], ids=[doc_id])
print("Document added.")
elif command == "query":
query_text = sys.argv[2]
user_id = sys.argv[3] # Not directly used by query, but can be used for filtering if implemented
results = db.query(query_text)
# Filter results by user_id if metadata contains it
filtered_results = {"documents": [[]], "metadatas": [[]], "distances": [[]], "ids": [[]]}
if results["metadatas"] and results["metadatas"][0]:
for i, meta in enumerate(results["metadatas"][0]):
if meta.get("userId") == user_id:
filtered_results["documents"][0].append(results["documents"][0][i])
filtered_results["metadatas"][0].append(results["metadatas"][0][i])
filtered_results["distances"][0].append(results["distances"][0][i])
filtered_results["ids"][0].append(results["ids"][0][i])
print(json.dumps(filtered_results))
elif command == "delete":
doc_id = sys.argv[2]
db.delete_documents(ids=[doc_id])
print("Document deleted.")
elif command == "count":
print(db.count())
elif command == "get_all":
results = db.get_all_documents()
print(json.dumps(results))
else:
print(f"Unknown command: {command}", file=sys.stderr)
sys.exit(1)