Underground-Digital/Workflow-Engine
0
1import logging2import time3 4from core.rag.datasource.retrieval_service import RetrievalService5from core.rag.models.document import Document6from core.rag.retrieval.retrieval_methods import RetrievalMethod7from extensions.ext_database import db8from models.account import Account9from models.dataset import Dataset, DatasetQuery, DocumentSegment10 11default_retrieval_model = {12 "search_method": RetrievalMethod.SEMANTIC_SEARCH.value,13 "reranking_enable": False,14 "reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},15 "top_k": 2,16 "score_threshold_enabled": False,17}18 19 20class HitTestingService:21 @classmethod22 def retrieve(23 cls,24 dataset: Dataset,25 query: str,26 account: Account,27 retrieval_model: dict,28 external_retrieval_model: dict,29 limit: int = 10,30 ) -> dict:31 if dataset.available_document_count == 0 or dataset.available_segment_count == 0:32 return {33 "query": {34 "content": query,35 "tsne_position": {"x": 0, "y": 0},36 },37 "records": [],38 }39 40 start = time.perf_counter()41 42 # get retrieval model , if the model is not setting , using default43 if not retrieval_model:44 retrieval_model = dataset.retrieval_model or default_retrieval_model45 46 all_documents = RetrievalService.retrieve(47 retrieval_method=retrieval_model.get("search_method", "semantic_search"),48 dataset_id=dataset.id,49 query=cls.escape_query_for_search(query),50 top_k=retrieval_model.get("top_k", 2),51 score_threshold=retrieval_model.get("score_threshold", 0.0)52 if retrieval_model["score_threshold_enabled"]53 else 0.0,54 reranking_model=retrieval_model.get("reranking_model", None)55 if retrieval_model["reranking_enable"]56 else None,57 reranking_mode=retrieval_model.get("reranking_mode") or "reranking_model",58 weights=retrieval_model.get("weights", None),59 )60 61 end = time.perf_counter()62 logging.debug(f"Hit testing retrieve in {end - start:0.4f} seconds")63 64 dataset_query = DatasetQuery(65 dataset_id=dataset.id, content=query, source="hit_testing", created_by_role="account", created_by=account.id66 )67 68 db.session.add(dataset_query)69 db.session.commit()70 71 return cls.compact_retrieve_response(dataset, query, all_documents)72 73 @classmethod74 def external_retrieve(75 cls,76 dataset: Dataset,77 query: str,78 account: Account,79 external_retrieval_model: dict,80 ) -> dict:81 if dataset.provider != "external":82 return {83 "query": {"content": query},84 "records": [],85 }86 87 start = time.perf_counter()88 89 all_documents = RetrievalService.external_retrieve(90 dataset_id=dataset.id,91 query=cls.escape_query_for_search(query),92 external_retrieval_model=external_retrieval_model,93 )94 95 end = time.perf_counter()96 logging.debug(f"External knowledge hit testing retrieve in {end - start:0.4f} seconds")97 98 dataset_query = DatasetQuery(99 dataset_id=dataset.id, content=query, source="hit_testing", created_by_role="account", created_by=account.id100 )101 102 db.session.add(dataset_query)103 db.session.commit()104 105 return cls.compact_external_retrieve_response(dataset, query, all_documents)106 107 @classmethod108 def compact_retrieve_response(cls, dataset: Dataset, query: str, documents: list[Document]):109 records = []110 111 for document in documents:112 index_node_id = document.metadata["doc_id"]113 114 segment = (115 db.session.query(DocumentSegment)116 .filter(117 DocumentSegment.dataset_id == dataset.id,118 DocumentSegment.enabled == True,119 DocumentSegment.status == "completed",120 DocumentSegment.index_node_id == index_node_id,121 )122 .first()123 )124 125 if not segment:126 continue127 128 record = {129 "segment": segment,130 "score": document.metadata.get("score", None),131 }132 133 records.append(record)134 135 return {136 "query": {137 "content": query,138 },139 "records": records,140 }141 142 @classmethod143 def compact_external_retrieve_response(cls, dataset: Dataset, query: str, documents: list):144 records = []145 if dataset.provider == "external":146 for document in documents:147 record = {148 "content": document.get("content", None),149 "title": document.get("title", None),150 "score": document.get("score", None),151 "metadata": document.get("metadata", None),152 }153 records.append(record)154 return {155 "query": {156 "content": query,157 },158 "records": records,159 }160 161 @classmethod162 def hit_testing_args_check(cls, args):163 query = args["query"]164 165 if not query or len(query) > 250:166 raise ValueError("Query is required and cannot exceed 250 characters")167 168 @staticmethod169 def escape_query_for_search(query: str) -> str:170 return query.replace('"', '\\"')171 