MD2204/multi_modality
1
1import os
2from typing import List, Optional
3from pinecone import Pinecone
4from langchain_pinecone import PineconeVectorStore
5from langchain_huggingface import HuggingFaceEmbeddings
6from langchain_core.documents import Document
7from src.config import (
8 EMBEDDING_MODEL_NAME,
9 EMBEDDING_DEVICE,
10 PINECONE_API_KEY,
11 PINECONE_INDEX_NAME
12)
13
14def get_embeddings():
15 """Initialize HuggingFace embeddings."""
16 return HuggingFaceEmbeddings(
17 model_name=EMBEDDING_MODEL_NAME,
18 model_kwargs={'device': EMBEDDING_DEVICE},
19 encode_kwargs={'normalize_embeddings': True}
20 )
21
22def _get_pinecone_index():
23 """Get a Pinecone index client for direct operations."""
24 pc = Pinecone(api_key=PINECONE_API_KEY)
25 return pc.Index(PINECONE_INDEX_NAME)
26
27def load_vector_store() -> Optional[PineconeVectorStore]:
28 """Connect to the existing Pinecone index."""
29 try:
30 embeddings = get_embeddings()
31 vectorstore = PineconeVectorStore(
32 index_name=PINECONE_INDEX_NAME,
33 embedding=embeddings,
34 pinecone_api_key=PINECONE_API_KEY
35 )
36 return vectorstore
37 except Exception as e:
38 print(f"⚠️ Error connecting to Pinecone: {e}")
39 return None
40
41def get_existing_sources(vectorstore: PineconeVectorStore) -> set:
42 """Extract unique source paths from the Pinecone index using direct query."""
43 unique_sources = set()
44 try:
45 index = _get_pinecone_index()
46 # Use list to get all vector IDs, then fetch their metadata
47 # For efficiency, we'll do a dummy query and check results
48 stats = index.describe_index_stats()
49 total_vectors = stats.get('total_vector_count', 0)
50
51 if total_vectors == 0:
52 return unique_sources
53
54 # Use a dummy query to fetch vectors with their metadata
55 embeddings = get_embeddings()
56 dummy_vector = embeddings.embed_query("dummy")
57
58 results = index.query(
59 vector=dummy_vector,
60 top_k=min(total_vectors, 10000),
61 include_metadata=True
62 )
63
64 for match in results.get('matches', []):
65 metadata = match.get('metadata', {})
66 source = metadata.get('source', '')
67 if source:
68 normalized_source = os.path.normpath(os.path.abspath(source))
69 unique_sources.add(normalized_source)
70
71 except Exception as e:
72 print(f"⚠️ Error getting existing sources: {e}")
73
74 return unique_sources
75
76def update_vector_store(documents: List[Document]) -> str:
77 """
78 Add new documents to the Pinecone vector store.
79 Skips documents that are already present based on their source path.
80 """
81 vectorstore = load_vector_store()
82
83 if not vectorstore:
84 msg = f"🆕 Creating vector store with {len(documents)} chunks."
85 print(msg)
86 embeddings = get_embeddings()
87 PineconeVectorStore.from_documents(
88 documents,
89 embedding=embeddings,
90 index_name=PINECONE_INDEX_NAME,
91 pinecone_api_key=PINECONE_API_KEY
92 )
93 return msg
94
95 existing_sources = get_existing_sources(vectorstore)
96
97 # Filter documents
98 new_documents = []
99 skipped_count = 0
100
101 for doc in documents:
102 source = doc.metadata.get('source')
103 if source:
104 normalized_source = os.path.normpath(os.path.abspath(source))
105 if normalized_source in existing_sources:
106 skipped_count += 1
107 continue
108
109 new_documents.append(doc)
110
111 if not new_documents:
112 msg = f"ℹ️ No new documents to add. Skipped {skipped_count} chunks from existing files."
113 print(msg)
114 return msg
115
116 msg = f"ℹ️ Adding {len(new_documents)} new chunks. Skipped {skipped_count} existing chunks."
117 print(msg)
118 vectorstore.add_documents(new_documents)
119
120 return msg
121
122def clear_vector_store() -> str:
123 """Delete all vectors from the Pinecone index for a fresh re-ingestion."""
124 try:
125 index = _get_pinecone_index()
126 index.delete(delete_all=True)
127 msg = "🗑️ Cleared all vectors from Pinecone index."
128 print(msg)
129 return msg
130 except Exception as e:
131 msg = f"❌ Error clearing Pinecone index: {e}"
132 print(msg)
133 return msg
134 