| import logging |
| import time |
|
|
| from core.rag.datasource.retrieval_service import RetrievalService |
| from core.rag.models.document import Document |
| from core.rag.retrieval.retrieval_methods import RetrievalMethod |
| from extensions.ext_database import db |
| from models.account import Account |
| from models.dataset import Dataset, DatasetQuery, DocumentSegment |
|
|
| default_retrieval_model = { |
| "search_method": RetrievalMethod.SEMANTIC_SEARCH.value, |
| "reranking_enable": False, |
| "reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""}, |
| "top_k": 2, |
| "score_threshold_enabled": False, |
| } |
|
|
|
|
| class HitTestingService: |
| @classmethod |
| def retrieve( |
| cls, |
| dataset: Dataset, |
| query: str, |
| account: Account, |
| retrieval_model: dict, |
| external_retrieval_model: dict, |
| limit: int = 10, |
| ) -> dict: |
| if dataset.available_document_count == 0 or dataset.available_segment_count == 0: |
| return { |
| "query": { |
| "content": query, |
| "tsne_position": {"x": 0, "y": 0}, |
| }, |
| "records": [], |
| } |
|
|
| start = time.perf_counter() |
|
|
| |
| if not retrieval_model: |
| retrieval_model = dataset.retrieval_model or default_retrieval_model |
|
|
| all_documents = RetrievalService.retrieve( |
| retrieval_method=retrieval_model.get("search_method", "semantic_search"), |
| dataset_id=dataset.id, |
| query=cls.escape_query_for_search(query), |
| top_k=retrieval_model.get("top_k", 2), |
| score_threshold=retrieval_model.get("score_threshold", 0.0) |
| if retrieval_model["score_threshold_enabled"] |
| else 0.0, |
| reranking_model=retrieval_model.get("reranking_model", None) |
| if retrieval_model["reranking_enable"] |
| else None, |
| reranking_mode=retrieval_model.get("reranking_mode") or "reranking_model", |
| weights=retrieval_model.get("weights", None), |
| ) |
|
|
| end = time.perf_counter() |
| logging.debug(f"Hit testing retrieve in {end - start:0.4f} seconds") |
|
|
| dataset_query = DatasetQuery( |
| dataset_id=dataset.id, content=query, source="hit_testing", created_by_role="account", created_by=account.id |
| ) |
|
|
| db.session.add(dataset_query) |
| db.session.commit() |
|
|
| return cls.compact_retrieve_response(dataset, query, all_documents) |
|
|
| @classmethod |
| def external_retrieve( |
| cls, |
| dataset: Dataset, |
| query: str, |
| account: Account, |
| external_retrieval_model: dict, |
| ) -> dict: |
| if dataset.provider != "external": |
| return { |
| "query": {"content": query}, |
| "records": [], |
| } |
|
|
| start = time.perf_counter() |
|
|
| all_documents = RetrievalService.external_retrieve( |
| dataset_id=dataset.id, |
| query=cls.escape_query_for_search(query), |
| external_retrieval_model=external_retrieval_model, |
| ) |
|
|
| end = time.perf_counter() |
| logging.debug(f"External knowledge hit testing retrieve in {end - start:0.4f} seconds") |
|
|
| dataset_query = DatasetQuery( |
| dataset_id=dataset.id, content=query, source="hit_testing", created_by_role="account", created_by=account.id |
| ) |
|
|
| db.session.add(dataset_query) |
| db.session.commit() |
|
|
| return cls.compact_external_retrieve_response(dataset, query, all_documents) |
|
|
| @classmethod |
| def compact_retrieve_response(cls, dataset: Dataset, query: str, documents: list[Document]): |
| records = [] |
|
|
| for document in documents: |
| index_node_id = document.metadata["doc_id"] |
|
|
| segment = ( |
| db.session.query(DocumentSegment) |
| .filter( |
| DocumentSegment.dataset_id == dataset.id, |
| DocumentSegment.enabled == True, |
| DocumentSegment.status == "completed", |
| DocumentSegment.index_node_id == index_node_id, |
| ) |
| .first() |
| ) |
|
|
| if not segment: |
| continue |
|
|
| record = { |
| "segment": segment, |
| "score": document.metadata.get("score", None), |
| } |
|
|
| records.append(record) |
|
|
| return { |
| "query": { |
| "content": query, |
| }, |
| "records": records, |
| } |
|
|
| @classmethod |
| def compact_external_retrieve_response(cls, dataset: Dataset, query: str, documents: list): |
| records = [] |
| if dataset.provider == "external": |
| for document in documents: |
| record = { |
| "content": document.get("content", None), |
| "title": document.get("title", None), |
| "score": document.get("score", None), |
| "metadata": document.get("metadata", None), |
| } |
| records.append(record) |
| return { |
| "query": { |
| "content": query, |
| }, |
| "records": records, |
| } |
|
|
| @classmethod |
| def hit_testing_args_check(cls, args): |
| query = args["query"] |
|
|
| if not query or len(query) > 250: |
| raise ValueError("Query is required and cannot exceed 250 characters") |
|
|
| @staticmethod |
| def escape_query_for_search(query: str) -> str: |
| return query.replace('"', '\\"') |
|
|