chore: extract retrival method literal values into enum (#5060)

This commit is contained in:
Bowen Liang 2024-06-19 16:05:27 +08:00 • committed by GitHub
commit c923684edd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 47 additions and 19 deletions

View file

@ -15,6 +15,7 @@ from core.model_runtime.model_providers.__base.large_language_model import Large
from core.rag.datasource.retrieval_service import RetrievalService
from core.rag.models.document import Document
from core.rag.rerank.rerank import RerankRunner
from core.rag.retrieval.retrival_methods import RetrievalMethod
from core.rag.retrieval.router.multi_dataset_function_call_router import FunctionCallMultiDatasetRouter
from core.rag.retrieval.router.multi_dataset_react_route import ReactMultiDatasetRouter
from core.tools.tool.dataset_retriever.dataset_multi_retriever_tool import DatasetMultiRetrieverTool
@ -25,7 +26,7 @@ from models.dataset import Dataset, DatasetQuery, DocumentSegment
from models.dataset import Document as DatasetDocument
default_retrieval_model = {
'search_method': 'semantic_search',
'search_method': RetrievalMethod.SEMANTIC_SEARCH,
'reranking_enable': False,
'reranking_model': {
'reranking_provider_name': '',
@ -419,7 +420,7 @@ class DatasetRetrieval:
if retrieve_config.retrieve_strategy == DatasetRetrieveConfigEntity.RetrieveStrategy.SINGLE:
# get retrieval model config
default_retrieval_model = {
'search_method': 'semantic_search',
'search_method': RetrievalMethod.SEMANTIC_SEARCH,
'reranking_enable': False,
'reranking_model': {
'reranking_provider_name': '',

View file

@ -0,0 +1,15 @@
from enum import Enum
class RetrievalMethod(str, Enum):
SEMANTIC_SEARCH = 'semantic_search'
FULL_TEXT_SEARCH = 'full_text_search'
HYBRID_SEARCH = 'hybrid_search'
@staticmethod
def is_support_semantic_search(retrieval_method: str) -> bool:
return retrieval_method in {RetrievalMethod.SEMANTIC_SEARCH, RetrievalMethod.HYBRID_SEARCH}
@staticmethod
def is_support_fulltext_search(retrieval_method: str) -> bool:
return retrieval_method in {RetrievalMethod.FULL_TEXT_SEARCH, RetrievalMethod.HYBRID_SEARCH}