Underground-Digital/Workflow-Engine
0
1import threading2from typing import Optional3 4from flask import Flask, current_app5 6from core.rag.data_post_processor.data_post_processor import DataPostProcessor7from core.rag.datasource.keyword.keyword_factory import Keyword8from core.rag.datasource.vdb.vector_factory import Vector9from core.rag.rerank.rerank_type import RerankMode10from core.rag.retrieval.retrieval_methods import RetrievalMethod11from extensions.ext_database import db12from models.dataset import Dataset13from services.external_knowledge_service import ExternalDatasetService14 15default_retrieval_model = {16 "search_method": RetrievalMethod.SEMANTIC_SEARCH.value,17 "reranking_enable": False,18 "reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},19 "top_k": 2,20 "score_threshold_enabled": False,21}22 23 24class RetrievalService:25 @classmethod26 def retrieve(27 cls,28 retrieval_method: str,29 dataset_id: str,30 query: str,31 top_k: int,32 score_threshold: Optional[float] = 0.0,33 reranking_model: Optional[dict] = None,34 reranking_mode: Optional[str] = "reranking_model",35 weights: Optional[dict] = None,36 ):37 if not query:38 return []39 dataset = db.session.query(Dataset).filter(Dataset.id == dataset_id).first()40 if not dataset:41 return []42 43 if not dataset or dataset.available_document_count == 0 or dataset.available_segment_count == 0:44 return []45 all_documents = []46 threads = []47 exceptions = []48 # retrieval_model source with keyword49 if retrieval_method == "keyword_search":50 keyword_thread = threading.Thread(51 target=RetrievalService.keyword_search,52 kwargs={53 "flask_app": current_app._get_current_object(),54 "dataset_id": dataset_id,55 "query": query,56 "top_k": top_k,57 "all_documents": all_documents,58 "exceptions": exceptions,59 },60 )61 threads.append(keyword_thread)62 keyword_thread.start()63 # retrieval_model source with semantic64 if RetrievalMethod.is_support_semantic_search(retrieval_method):65 embedding_thread = threading.Thread(66 target=RetrievalService.embedding_search,67 kwargs={68 "flask_app": current_app._get_current_object(),69 "dataset_id": dataset_id,70 "query": query,71 "top_k": top_k,72 "score_threshold": score_threshold,73 "reranking_model": reranking_model,74 "all_documents": all_documents,75 "retrieval_method": retrieval_method,76 "exceptions": exceptions,77 },78 )79 threads.append(embedding_thread)80 embedding_thread.start()81 82 # retrieval source with full text83 if RetrievalMethod.is_support_fulltext_search(retrieval_method):84 full_text_index_thread = threading.Thread(85 target=RetrievalService.full_text_index_search,86 kwargs={87 "flask_app": current_app._get_current_object(),88 "dataset_id": dataset_id,89 "query": query,90 "retrieval_method": retrieval_method,91 "score_threshold": score_threshold,92 "top_k": top_k,93 "reranking_model": reranking_model,94 "all_documents": all_documents,95 "exceptions": exceptions,96 },97 )98 threads.append(full_text_index_thread)99 full_text_index_thread.start()100 101 for thread in threads:102 thread.join()103 104 if exceptions:105 exception_message = ";\n".join(exceptions)106 raise Exception(exception_message)107 108 if retrieval_method == RetrievalMethod.HYBRID_SEARCH.value:109 data_post_processor = DataPostProcessor(110 str(dataset.tenant_id), reranking_mode, reranking_model, weights, False111 )112 all_documents = data_post_processor.invoke(113 query=query, documents=all_documents, score_threshold=score_threshold, top_n=top_k114 )115 return all_documents116 117 @classmethod118 def external_retrieve(cls, dataset_id: str, query: str, external_retrieval_model: Optional[dict] = None):119 dataset = db.session.query(Dataset).filter(Dataset.id == dataset_id).first()120 if not dataset:121 return []122 all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval(123 dataset.tenant_id, dataset_id, query, external_retrieval_model124 )125 return all_documents126 127 @classmethod128 def keyword_search(129 cls, flask_app: Flask, dataset_id: str, query: str, top_k: int, all_documents: list, exceptions: list130 ):131 with flask_app.app_context():132 try:133 dataset = db.session.query(Dataset).filter(Dataset.id == dataset_id).first()134 135 keyword = Keyword(dataset=dataset)136 137 documents = keyword.search(cls.escape_query_for_search(query), top_k=top_k)138 all_documents.extend(documents)139 except Exception as e:140 exceptions.append(str(e))141 142 @classmethod143 def embedding_search(144 cls,145 flask_app: Flask,146 dataset_id: str,147 query: str,148 top_k: int,149 score_threshold: Optional[float],150 reranking_model: Optional[dict],151 all_documents: list,152 retrieval_method: str,153 exceptions: list,154 ):155 with flask_app.app_context():156 try:157 dataset = db.session.query(Dataset).filter(Dataset.id == dataset_id).first()158 159 vector = Vector(dataset=dataset)160 161 documents = vector.search_by_vector(162 cls.escape_query_for_search(query),163 search_type="similarity_score_threshold",164 top_k=top_k,165 score_threshold=score_threshold,166 filter={"group_id": [dataset.id]},167 )168 169 if documents:170 if (171 reranking_model172 and reranking_model.get("reranking_model_name")173 and reranking_model.get("reranking_provider_name")174 and retrieval_method == RetrievalMethod.SEMANTIC_SEARCH.value175 ):176 data_post_processor = DataPostProcessor(177 str(dataset.tenant_id), RerankMode.RERANKING_MODEL.value, reranking_model, None, False178 )179 all_documents.extend(180 data_post_processor.invoke(181 query=query, documents=documents, score_threshold=score_threshold, top_n=len(documents)182 )183 )184 else:185 all_documents.extend(documents)186 except Exception as e:187 exceptions.append(str(e))188 189 @classmethod190 def full_text_index_search(191 cls,192 flask_app: Flask,193 dataset_id: str,194 query: str,195 top_k: int,196 score_threshold: Optional[float],197 reranking_model: Optional[dict],198 all_documents: list,199 retrieval_method: str,200 exceptions: list,201 ):202 with flask_app.app_context():203 try:204 dataset = db.session.query(Dataset).filter(Dataset.id == dataset_id).first()205 206 vector_processor = Vector(207 dataset=dataset,208 )209 210 documents = vector_processor.search_by_full_text(cls.escape_query_for_search(query), top_k=top_k)211 if documents:212 if (213 reranking_model214 and reranking_model.get("reranking_model_name")215 and reranking_model.get("reranking_provider_name")216 and retrieval_method == RetrievalMethod.FULL_TEXT_SEARCH.value217 ):218 data_post_processor = DataPostProcessor(219 str(dataset.tenant_id), RerankMode.RERANKING_MODEL.value, reranking_model, None, False220 )221 all_documents.extend(222 data_post_processor.invoke(223 query=query, documents=documents, score_threshold=score_threshold, top_n=len(documents)224 )225 )226 else:227 all_documents.extend(documents)228 except Exception as e:229 exceptions.append(str(e))230 231 @staticmethod232 def escape_query_for_search(query: str) -> str:233 return query.replace('"', '\\"')234 