File size: 4,824 Bytes
c13f797 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 | 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)
|