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)