diff --git a/api/.env.example b/api/.env.example index 998d8b5d69a..77ec7ed139e 100644 --- a/api/.env.example +++ b/api/.env.example @@ -326,6 +326,7 @@ TIDB_VECTOR_PORT=4000 TIDB_VECTOR_USER=xxx.root TIDB_VECTOR_PASSWORD=xxxxxx TIDB_VECTOR_DATABASE=dify +TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH=false # Tidb on qdrant configuration TIDB_ON_QDRANT_URL=http://127.0.0.1 diff --git a/api/configs/middleware/vdb/tidb_vector_config.py b/api/configs/middleware/vdb/tidb_vector_config.py index 0ebf226bea6..172b5a7e56e 100644 --- a/api/configs/middleware/vdb/tidb_vector_config.py +++ b/api/configs/middleware/vdb/tidb_vector_config.py @@ -31,3 +31,8 @@ class TiDBVectorConfig(BaseSettings): description="Name of the TiDB Vector database to connect to", default=None, ) + + TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH: bool = Field( + description="Enable TiDB Vector full-text and hybrid search features", + default=False, + ) diff --git a/api/controllers/console/datasets/datasets.py b/api/controllers/console/datasets/datasets.py index 19a6c4dc0dc..9d444978c01 100644 --- a/api/controllers/console/datasets/datasets.py +++ b/api/controllers/console/datasets/datasets.py @@ -360,7 +360,6 @@ def _get_retrieval_methods_by_vector_type(vector_type: str | None, is_mock: bool # Define vector database types that only support semantic search semantic_only_types = { VectorType.RELYT, - VectorType.TIDB_VECTOR, VectorType.CHROMA, VectorType.PGVECTO_RS, VectorType.VIKINGDB, @@ -408,6 +407,9 @@ def _get_retrieval_methods_by_vector_type(vector_type: str | None, is_mock: bool if vector_type == VectorType.MILVUS: return semantic_methods if is_mock else full_methods + if vector_type == VectorType.TIDB_VECTOR: + return full_methods if dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH else semantic_methods + if vector_type in semantic_only_types: return semantic_methods elif vector_type in full_search_types: diff --git a/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py b/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py index 9f80ae5a76f..1fadb5aef58 100644 --- a/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py +++ b/api/providers/vdb/vdb-tidb-vector/src/dify_vdb_tidb_vector/tidb_vector.py @@ -20,6 +20,8 @@ from models.dataset import Dataset logger = logging.getLogger(__name__) +FULLTEXT_INDEX_NAME = "idx_text" + class TiDBVectorConfig(BaseModel): host: str @@ -28,6 +30,7 @@ class TiDBVectorConfig(BaseModel): password: str database: str program_name: str + enable_fulltext_search: bool = False @model_validator(mode="before") @classmethod @@ -95,10 +98,15 @@ class TiDBVector(BaseVector): logger.info("_create_collection, collection_name %s", self._collection_name) lock_name = f"vector_indexing_lock_{self._collection_name}" with redis_client.lock(lock_name, timeout=20): - collection_exist_cache_key = f"vector_indexing_{self._collection_name}" + collection_exist_cache_key = self._collection_exist_cache_key() if redis_client.get(collection_exist_cache_key): return tidb_dist_func = self._get_distance_func() + fulltext_index_statement = ( + f",\n FULLTEXT INDEX {FULLTEXT_INDEX_NAME} (text) WITH PARSER MULTILINGUAL" + if self._client_config.enable_fulltext_search + else "" + ) with sessionmaker(bind=self._engine).begin() as session: create_statement = sql_text(f""" CREATE TABLE IF NOT EXISTS {self._collection_name} ( @@ -113,11 +121,52 @@ class TiDBVector(BaseVector): KEY (doc_id), KEY (document_id), VECTOR INDEX idx_vector (({tidb_dist_func}(vector))) USING HNSW + {fulltext_index_statement} ); """) session.execute(create_statement) + if self._client_config.enable_fulltext_search: + self._ensure_fulltext_index(session) redis_client.set(collection_exist_cache_key, 1, ex=3600) + def _collection_exist_cache_key(self) -> str: + search_mode = "fulltext" if self._client_config.enable_fulltext_search else "semantic" + return f"vector_indexing_{self._collection_name}_{search_mode}" + + def _ensure_fulltext_index(self, session) -> None: + index_check_statement = sql_text(""" + SELECT COUNT(1) + FROM INFORMATION_SCHEMA.STATISTICS + WHERE TABLE_SCHEMA = DATABASE() + AND TABLE_NAME = :table_name + AND INDEX_NAME = :index_name + """) + result = session.execute( + index_check_statement, + params={ + "index_name": FULLTEXT_INDEX_NAME, + "table_name": self._collection_name, + }, + ) + if result.scalar(): + return + + session.execute( + sql_text( + f"ALTER TABLE {self._collection_name} " + f"ADD FULLTEXT INDEX {FULLTEXT_INDEX_NAME} (text) WITH PARSER MULTILINGUAL;" + ) + ) + + @staticmethod + def _document_ids_filter_condition(document_ids_filter: list[str] | None) -> tuple[str, dict[str, str]]: + if not document_ids_filter: + return "", {} + + filter_params = {f"document_id_{index}": document_id for index, document_id in enumerate(document_ids_filter)} + placeholders = ", ".join(f":{param_name}" for param_name in filter_params) + return f"document_id IN ({placeholders})", filter_params + @override def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs): table = self._table(len(embeddings[0])) @@ -127,8 +176,8 @@ class TiDBVector(BaseVector): chunks_table_data = [] with self._engine.connect() as conn, conn.begin(): - for id, text, meta, embedding in zip(ids, texts, metas, embeddings): - chunks_table_data.append({"id": id, "vector": embedding, "text": text, "meta": meta}) + for doc_id, text, meta, embedding in zip(ids, texts, metas, embeddings): + chunks_table_data.append({"id": doc_id, "vector": embedding, "text": text, "meta": meta}) # Execute the batch insert when the batch size is reached if len(chunks_table_data) == 500: @@ -205,9 +254,10 @@ class TiDBVector(BaseVector): tidb_dist_func = self._get_distance_func() document_ids_filter = kwargs.get("document_ids_filter") where_clause = "" + filter_params: dict[str, str] = {} if document_ids_filter: - document_ids = ", ".join(f"'{id}'" for id in document_ids_filter) - where_clause = f" WHERE meta->>'$.document_id' in ({document_ids}) " + document_ids_filter_condition, filter_params = self._document_ids_filter_condition(document_ids_filter) + where_clause = f" WHERE {document_ids_filter_condition} " with Session(self._engine) as session: select_statement = sql_text(f""" @@ -230,6 +280,7 @@ class TiDBVector(BaseVector): "query_vector_str": query_vector_str, "distance": distance, "top_k": top_k, + **filter_params, }, ) results = [(row[0], row[1], row[2]) for row in res] @@ -241,8 +292,51 @@ class TiDBVector(BaseVector): @override def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]: - # tidb doesn't support bm25 search - return [] + if not self._client_config.enable_fulltext_search or not query: + return [] + + top_k = kwargs.get("top_k", 4) + score_threshold = float(kwargs.get("score_threshold") or 0.0) + document_ids_filter = kwargs.get("document_ids_filter") + + where_conditions = ["FTS_MATCH_WORD(text, :query)"] + filter_params: dict[str, str] = {} + if document_ids_filter: + document_ids_filter_condition, filter_params = self._document_ids_filter_condition(document_ids_filter) + where_conditions.append(document_ids_filter_condition) + where_clause = " AND ".join(where_conditions) + + docs = [] + with Session(self._engine) as session: + select_statement = sql_text(f""" + SELECT meta, text, score + FROM ( + SELECT + meta, + text, + FTS_MATCH_WORD(text, :query) AS score + FROM {self._collection_name} + WHERE {where_clause} + ORDER BY score DESC + LIMIT :top_k + ) t + WHERE score >= :score_threshold + """) + res = session.execute( + select_statement, + params={ + "query": query, + "score_threshold": score_threshold, + "top_k": top_k, + **filter_params, + }, + ) + results = [(row[0], row[1], row[2]) for row in res] + for meta, text, score in results: + metadata = parse_metadata_json(meta) + metadata["score"] = score + docs.append(Document(page_content=text, metadata=metadata)) + return docs @override def delete(self): @@ -280,5 +374,6 @@ class TiDBVectorFactory(AbstractVectorFactory): password=dify_config.TIDB_VECTOR_PASSWORD or "", database=dify_config.TIDB_VECTOR_DATABASE or "", program_name=dify_config.APPLICATION_NAME, + enable_fulltext_search=dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH, ), ) diff --git a/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py b/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py index ed03cbee88d..47d63f0641d 100644 --- a/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py +++ b/api/providers/vdb/vdb-tidb-vector/tests/unit_tests/test_tidb_vector.py @@ -25,6 +25,7 @@ def _config(tidb_module): password="secret", database="dify", program_name="dify-app", + enable_fulltext_search=False, ) @@ -118,11 +119,13 @@ def test_create_collection_skips_when_cache_hit(tidb_module, monkeypatch: pytest vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) vector._collection_name = "collection_1" vector._engine = MagicMock() + vector._client_config = _config(tidb_module) tidb_module.Session = MagicMock() vector._create_collection(3) + tidb_module.redis_client.get.assert_called_once_with("vector_indexing_collection_1_semantic") tidb_module.Session.assert_not_called() tidb_module.redis_client.set.assert_not_called() @@ -151,15 +154,90 @@ def test_create_collection_executes_create_sql_and_sets_cache(tidb_module, monke vector._collection_name = "collection_1" vector._engine = MagicMock() vector._distance_func = "l2" + vector._client_config = _config(tidb_module) vector._create_collection(3) sql = str(session.execute.call_args.args[0]) assert "VECTOR(3)" in sql assert "VEC_L2_DISTANCE" in sql + assert "FULLTEXT INDEX" not in sql tidb_module.redis_client.set.assert_called_once() +def test_create_collection_adds_fulltext_index_when_enabled(tidb_module, monkeypatch: pytest.MonkeyPatch): + lock = MagicMock() + lock.__enter__.return_value = None + lock.__exit__.return_value = None + monkeypatch.setattr(tidb_module.redis_client, "lock", MagicMock(return_value=lock)) + monkeypatch.setattr(tidb_module.redis_client, "get", MagicMock(return_value=None)) + monkeypatch.setattr(tidb_module.redis_client, "set", MagicMock()) + + session = MagicMock() + + class _BeginCtx: + def __enter__(self): + return session + + def __exit__(self, exc_type, exc, tb): + return False + + mock_sm = MagicMock(begin=MagicMock(return_value=_BeginCtx())) + monkeypatch.setattr(tidb_module, "sessionmaker", lambda **kwargs: mock_sm) + + vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) + vector._collection_name = "collection_1" + vector._engine = MagicMock() + vector._distance_func = "cosine" + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + + vector._create_collection(3) + + sql = str(session.execute.call_args_list[0].args[0]) + assert "FULLTEXT INDEX idx_text (text) WITH PARSER MULTILINGUAL" in sql + + +def test_create_collection_ensures_fulltext_index_when_enabled(tidb_module, monkeypatch: pytest.MonkeyPatch): + lock = MagicMock() + lock.__enter__.return_value = None + lock.__exit__.return_value = None + monkeypatch.setattr(tidb_module.redis_client, "lock", MagicMock(return_value=lock)) + monkeypatch.setattr(tidb_module.redis_client, "get", MagicMock(return_value=None)) + monkeypatch.setattr(tidb_module.redis_client, "set", MagicMock()) + + session = MagicMock() + index_check_result = MagicMock() + index_check_result.scalar.return_value = 0 + session.execute.side_effect = [None, index_check_result, None] + + class _BeginCtx: + def __enter__(self): + return session + + def __exit__(self, exc_type, exc, tb): + return False + + mock_sm = MagicMock(begin=MagicMock(return_value=_BeginCtx())) + monkeypatch.setattr(tidb_module, "sessionmaker", lambda **kwargs: mock_sm) + + vector = tidb_module.TiDBVector.__new__(tidb_module.TiDBVector) + vector._collection_name = "collection_1" + vector._engine = MagicMock() + vector._distance_func = "cosine" + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + + vector._create_collection(3) + + executed_sql = [str(call.args[0]) for call in session.execute.call_args_list] + assert "INFORMATION_SCHEMA.STATISTICS" in executed_sql[1] + assert "ALTER TABLE collection_1 ADD FULLTEXT INDEX idx_text (text) WITH PARSER MULTILINGUAL" in executed_sql[2] + assert session.execute.call_args_list[1].kwargs["params"] == { + "index_name": "idx_text", + "table_name": "collection_1", + } + tidb_module.redis_client.get.assert_called_once_with("vector_indexing_collection_1_fulltext") + + def test_add_texts_batches_inserts_and_returns_ids(tidb_module, monkeypatch: pytest.MonkeyPatch): class _InsertStmt: def __init__(self, table): @@ -215,10 +293,45 @@ def tidb_vector_with_session(tidb_module, monkeypatch: pytest.MonkeyPatch): return vector, session, tidb_module -# 1. search_by_full_text returns empty -def test_search_by_full_text_returns_empty(tidb_vector_with_session): - vector, _, _ = tidb_vector_with_session +# 1. search_by_full_text returns empty when disabled +def test_search_by_full_text_returns_empty_when_disabled(tidb_vector_with_session): + vector, session, tidb_module = tidb_vector_with_session + vector._client_config = _config(tidb_module) assert vector.search_by_full_text("query") == [] + session.execute.assert_not_called() + + +def test_search_by_full_text_queries_tidb_fts_and_scores(tidb_vector_with_session): + vector, session, tidb_module = tidb_vector_with_session + vector._client_config = _config(tidb_module).model_copy(update={"enable_fulltext_search": True}) + session.execute.return_value = [ + ('{"doc_id":"id-1","document_id":"d-1"}', "text-1", 0.8), + ('{"doc_id":"id-2","document_id":"d-2"}', "text-2", 0.6), + ] + + docs = vector.search_by_full_text( + "search query", + top_k=2, + score_threshold=0.5, + document_ids_filter=["d-1", "d'2"], + ) + + assert len(docs) == 2 + assert docs[0].page_content == "text-1" + assert docs[0].metadata["score"] == pytest.approx(0.8) + assert docs[1].metadata["score"] == pytest.approx(0.6) + sql = str(session.execute.call_args.args[0]) + params = session.execute.call_args.kwargs["params"] + assert "FTS_MATCH_WORD(text, :query)" in sql + assert "document_id IN (:document_id_0, :document_id_1)" in sql + assert "d'2" not in sql + assert params == { + "document_id_0": "d-1", + "document_id_1": "d'2", + "query": "search query", + "score_threshold": 0.5, + "top_k": 2, + } # 2. text_exists returns True when ids found @@ -378,15 +491,18 @@ def test_search_by_vector_filters_and_scores(tidb_module, monkeypatch: pytest.Mo [0.1, 0.2], top_k=2, score_threshold=0.5, - document_ids_filter=["d-1", "d-2"], + document_ids_filter=["d-1", "d'2"], ) assert len(docs) == 2 assert docs[0].metadata["score"] == pytest.approx(0.8) assert docs[1].metadata["score"] == pytest.approx(0.6) sql = str(session.execute.call_args.args[0]) params = session.execute.call_args.kwargs["params"] - assert "meta->>'$.document_id' in ('d-1', 'd-2')" in sql + assert "document_id IN (:document_id_0, :document_id_1)" in sql + assert "d'2" not in sql assert params["distance"] == pytest.approx(0.5) + assert params["document_id_0"] == "d-1" + assert params["document_id_1"] == "d'2" assert params["top_k"] == 2 session.commit.assert_not_called() @@ -428,6 +544,7 @@ def test_tidb_factory_uses_existing_or_generated_collection(tidb_module, monkeyp monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_USER", "root") monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_PASSWORD", "secret") monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_DATABASE", "dify") + monkeypatch.setattr(tidb_module.dify_config, "TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", True) monkeypatch.setattr(tidb_module.dify_config, "APPLICATION_NAME", "dify-app") with patch.object(tidb_module, "TiDBVector", return_value="vector") as vector_cls: @@ -438,4 +555,5 @@ def test_tidb_factory_uses_existing_or_generated_collection(tidb_module, monkeyp assert result_2 == "vector" assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection" assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection" + assert vector_cls.call_args_list[0].kwargs["config"].enable_fulltext_search is True assert dataset_without_index.index_struct is not None diff --git a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py index 68b1ad0dd60..72c095e2356 100644 --- a/api/tests/unit_tests/controllers/console/datasets/test_datasets.py +++ b/api/tests/unit_tests/controllers/console/datasets/test_datasets.py @@ -29,12 +29,15 @@ from controllers.console.datasets.datasets import ( DatasetRetrievalSettingApi, DatasetRetrievalSettingMockApi, DatasetUseCheckApi, + _get_retrieval_methods_by_vector_type, ) from controllers.console.datasets.error import DatasetInUseError, DatasetNameDuplicateError, IndexingEstimateError from core.entities.knowledge_entities import IndexingEstimate from core.errors.error import LLMBadRequestError, ProviderTokenNotInitError from core.provider_manager import ProviderManager +from core.rag.datasource.vdb.vector_type import VectorType from core.rag.index_processor.constant.index_type import IndexStructureType +from core.rag.retrieval.retrieval_methods import RetrievalMethod from extensions.storage.storage_type import StorageType from models.account import Account, TenantAccountRole from models.dataset import Dataset, DatasetQuery, Document @@ -1427,6 +1430,28 @@ class TestDatasetRetrievalSettingApi: response = method(api) assert "retrieval_method" in response + def test_tidb_vector_returns_semantic_only_when_fulltext_disabled(self): + with patch( + "controllers.console.datasets.datasets.dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", + False, + ): + response = _get_retrieval_methods_by_vector_type(VectorType.TIDB_VECTOR) + + assert response["retrieval_method"] == [RetrievalMethod.SEMANTIC_SEARCH.value] + + def test_tidb_vector_returns_full_methods_when_fulltext_enabled(self): + with patch( + "controllers.console.datasets.datasets.dify_config.TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH", + True, + ): + response = _get_retrieval_methods_by_vector_type(VectorType.TIDB_VECTOR) + + assert response["retrieval_method"] == [ + RetrievalMethod.SEMANTIC_SEARCH.value, + RetrievalMethod.FULL_TEXT_SEARCH.value, + RetrievalMethod.HYBRID_SEARCH.value, + ] + class TestDatasetRetrievalSettingMockApi: def test_get_success(self, app: Flask): diff --git a/docker/envs/core-services/shared.env.example b/docker/envs/core-services/shared.env.example index be9f8312cdf..c3843018b0c 100644 --- a/docker/envs/core-services/shared.env.example +++ b/docker/envs/core-services/shared.env.example @@ -417,6 +417,7 @@ TIDB_VECTOR_HOST=tidb TIDB_VECTOR_PORT=4000 TIDB_VECTOR_USER= TIDB_VECTOR_PASSWORD= +TIDB_VECTOR_ENABLE_FULLTEXT_SEARCH=false TIDB_ON_QDRANT_CLIENT_TIMEOUT=20 TIDB_ON_QDRANT_GRPC_ENABLED=false TIDB_ON_QDRANT_GRPC_PORT=6334