Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
retrieval_service.py234 linesDownload Raw Back to datasource
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