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)