Download src/memory/vector_db.py from agenthinkmesh/sane-assistant: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/agenthinkmesh/sane-assistant/resolve/main/src/memory/vector_db.py
- Command line
-
hf download hf://agenthinkmesh/sane-assistant/src/memory/vector_db.py
-
curl -L -o vector_db.py https://huggingface.co/agenthinkmesh/sane-assistant/resolve/main/src/memory/vector_db.py
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) | |