Spaces:
Sleeping
Sleeping
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
|