Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
hit_testing_service.py171 linesDownload Raw Back to services
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