File size: 6,183 Bytes
338036b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
from .BaseController import BaseController
from models.db_schemes.minirag.schemes import Project, DataChunk
from stores.LLM.LLMEnums import DocumentTypeEnum
from typing import List
import logging

logger = logging.getLogger(__name__)

class NLPController(BaseController):
     def __init__(self,vectordb_client, generation_client,embedding_client, template_parser):
          super().__init__()
          self.vectordb_client = vectordb_client  # Note: VectorDBProvider functions are async so we need to await them when calling them.
          self.embedding_client = embedding_client
          self.generation_client = generation_client
          self.template_parser = template_parser

     def create_collection_name(self, project_id):
          return f"collection_{self.vectordb_client.default_vector_size}_{project_id}".strip() 
     
     async def reset_vector_db_collection(self, project: Project): 
          # Bug: this used `Project.project_id` (the class), which has no runtime value.
          # Fix: read `project.project_id` from the actual project instance passed in.
          collection_name = self.create_collection_name(project_id=project.project_id)
          return await self.vectordb_client.delete_collection(collection_name= collection_name)
     
     async def get_vector_db_collection_info(self, project: Project):
          # Same class-vs-instance bug here: the collection name must come from the loaded project record.
          collection_name = self.create_collection_name(project_id=project.project_id)
          collection_info = await self.vectordb_client.get_collection_info(collection_name=collection_name)
          return collection_info
 
     async def index_into_vector_db(self, project:Project, chunks: List[DataChunk], 
                              chunks_ids : List[int]):
          # Bug: using `Project.project_id` raised `AttributeError: project_id`.
          # Fix: derive the collection name from the current project instance.
          collection_name = self.create_collection_name(project_id=project.project_id)

          # 2. manage items 
          texts = [c.chunk_text for c in chunks]
          metadata = [c.chunk_metadata for c in chunks]
          # Bug: embedding one chunk per API call quickly hits Cohere's trial limit.
          # Fix: send the whole page of chunk texts in a single batched embed request.
          vectors = self.embedding_client.embed_text (
               text=texts,
               document_type=DocumentTypeEnum.DOCUMENT.value
          )
          if not vectors or len(vectors) != len(texts):
               return False

          # The route prepares the collection once before paging starts.
          # Recreating it here would reset the table on every batch when `do_reset=True`.
          return await self.vectordb_client.insert_many(
               collection_name=collection_name,
               texts=texts,
               vectors=vectors,
               metadata=metadata,
               recored_ids = chunks_ids,
          )
     
     async def search_vector_db_collection(self, project: Project, text: str, limit: int=10):
          query_vector = None
          collection_name = self.create_collection_name(project_id=project.project_id)
          vectors = self.embedding_client.embed_text(text = text,
                                             document_type = DocumentTypeEnum.QUERY.value)
          
          if not vectors or len(vectors) == 0 :
               return False
          
          if isinstance(vectors, list) or len(vectors) > 0:
               query_vector = vectors[0]
          
          if not query_vector:
               logger.error("Failed to get query vector")
               return False

          results = await self.vectordb_client.search_by_vector(
               collection_name = collection_name,
               vector = query_vector,
               limit = limit 
          )
          if not results:
               return 
          
          return [
               result.dict()        
               for result in results 
          ]
      
     
     async def answer_rag_question(self, project: Project, query: str, limit: int = 10):
          # step1: retrieve related documents
          retieved_documents = await self.search_vector_db_collection(project=project, text=query, limit=limit)

          # Bug: returning None here caused a TypeError when the route tried to unpack
          # the result as (answer, full_prompt, chat_history).
          # Fix: always return a consistent tuple so the route can handle it cleanly.
          if not retieved_documents or len(retieved_documents) == 0:
               logger.warning("No documents retrieved from vector DB for query: '%s'", query)
               return None, None, None

          logger.info("Retrieved %d documents from vector DB.", len(retieved_documents))

          # step2: construct LLM prompt
          system_prompt = self.template_parser.get("rag", "system_prompt")
          print(retieved_documents)
     
          documents_prompts = "\n".join([
               self.template_parser.get("rag", "document_prompt", {
                    "doc_num": idx,
                    "chunk_text": self.generation_client.process_text(doc["text"]) 
               })
               for idx, doc in enumerate(retieved_documents)
          ])

          footer_prompt = self.template_parser.get("rag", "footer_template", {
               "query": query
          })

          chat_history = [
               self.generation_client.construct_prompt(
                    prompt=system_prompt,
                    role=self.generation_client.enums.SYSTEM
               )
          ]

          full_prompt = "\n\n".join([documents_prompts, footer_prompt])

          logger.info("Sending prompt to LLM (length: %d chars).", len(full_prompt))

          answer = self.generation_client.generate_text(
               prompt=full_prompt,
               chat_history=chat_history
          )

          if not answer:
               logger.error("LLM returned no answer. Check that Ollama is running on port 11434.")

          return answer, full_prompt, chat_history