From 75bff4bcd92853d99bfda04bfbb5107a06722133 Mon Sep 17 00:00:00 2001 From: Classic298 <27028174+Classic298@users.noreply.github.com> Date: Wed, 30 Sep 2026 19:03:47 +0200 Subject: [PATCH] feat: native hybrid search for Milvus and Milvus multitenancy (#31645) With hybrid search on, Milvus installs fetch every chunk of a collection and score BM25 in Python for each search. In both Milvus modes, one collection per knowledge base and multitenancy, Milvus now runs the keyword half itself with its built-in BM25 full-text search and merges it with the vector results, as pgvector already does. Milvus cannot add a BM25 index to an existing collection, so each collection gets a second, text-only collection next to it; existing installs build these once during startup when ENABLE_DB_MIGRATIONS is on, which copies the chunk text (extra storage roughly the size of that text) and leaves the original vectors and indexes untouched. On one standalone Milvus server the copy ran at about 11,000 chunks per second, around 40 minutes for 200 GB with 1536-dimension embeddings. Collections that cannot be copied, and Milvus servers older than 2.5, which have no BM25, keep using the existing hybrid search. Fixes #26243 --- .../open_webui/retrieval/vector/dbs/milvus.py | 252 +++++++++++++++++- .../vector/dbs/milvus_multitenancy.py | 140 +++++++++- 2 files changed, 381 insertions(+), 11 deletions(-) diff --git a/backend/open_webui/retrieval/vector/dbs/milvus.py b/backend/open_webui/retrieval/vector/dbs/milvus.py index d09f01c16d..d5aab4a1a8 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus.py @@ -4,6 +4,8 @@ NOTE: This vector database integration is community-supported and maintained on import logging import re +from collections.abc import Callable, Iterator +from concurrent.futures import ThreadPoolExecutor from typing import Any, Optional from open_webui.config import ( @@ -18,16 +20,18 @@ from open_webui.config import ( MILVUS_TOKEN, MILVUS_URI, ) +from open_webui.env import ENABLE_DB_MIGRATIONS from open_webui.retrieval.vector.main import ( GetResult, SearchResult, VectorDBBase, VectorItem, ) -from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata +from open_webui.retrieval.vector.utils import iter_filter_conditions, merge_hybrid_search_results, process_metadata from open_webui.utils.json_codec import JSONCodec -from pymilvus import DataType +from pymilvus import DataType, Function, FunctionType from pymilvus import MilvusClient as Client +from pymilvus.client.types import LoadState from pymilvus.exceptions import MilvusException log = logging.getLogger(__name__) @@ -36,6 +40,14 @@ log = logging.getLogger(__name__) # field). Clamp long chunks before insert so one oversized chunk can't fail the # whole batch and leave the file with zero embeddings. MILVUS_TEXT_MAX_LENGTH = 65535 +# Milvus cannot add a BM25 field to an existing collection, so BM25 search lives in a companion collection. +BM25_COLLECTION_SUFFIX = '_bm25' +BM25_STAGING_SUFFIX = '_bm25_staging' +# Milvus rejects messages over 64 MB by default: reads step down on failure, inserts stay well below. +BM25_BACKFILL_BATCH_SIZES = (16384, 4096, 1024, 256) +BM25_BACKFILL_MAX_BYTES = 16 * 1024 * 1024 +BM25_BACKFILL_WORKERS = 8 +BM25_BACKFILL_LOAD_TIMEOUT = 300 _SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$') @@ -68,6 +80,97 @@ def _metadata_exprs(filter: Optional[dict]) -> list[str]: return exprs +def _bm25_rows(rows: list[dict]) -> list[dict]: + bm25_rows = [] + for row in rows: + data = row.get('data') + text = data.get('text') if isinstance(data, dict) else None + bm25_rows.append( + {'id': row['id'], 'text': text if isinstance(text, str) else '', 'metadata': row.get('metadata') or {}} + ) + return bm25_rows + + +def _size_capped_batches(rows: list[dict]) -> Iterator[list[dict]]: + batch, batch_bytes = [], 0 + for row in rows: + row_bytes = len(JSONCodec.dumps(row).encode()) + if batch and batch_bytes + row_bytes > BM25_BACKFILL_MAX_BYTES: + yield batch + batch, batch_bytes = [], 0 + batch.append(row) + batch_bytes += row_bytes + if batch: + yield batch + + +def _update_bm25_collection(client: Client, bm25_collection: str, operation: str, **kwargs: Any): + # BM25 rows only feed hybrid search, so a failed BM25 write never fails the original one. + try: + if not client.has_collection(bm25_collection): + return + if operation == 'delete': + client.load_collection(bm25_collection) + getattr(client, operation)(collection_name=bm25_collection, **kwargs) + except MilvusException as e: + log.warning('Could not update BM25 collection %s: %s', bm25_collection, e) + + +def _backfill_bm25_collection( + client: Client, + collection: str, + create_bm25_collection: Callable[[str], None], + output_fields: list[str], + to_bm25_rows: Callable[[list[dict]], list[dict]], +): + # Built under a staging name so a crash never leaves a partial BM25 collection in use. + staging_collection = f'{collection}{BM25_STAGING_SUFFIX}' + log.info('Building BM25 collection for %s.', collection) + try: + client.drop_collection(staging_collection) + create_bm25_collection(staging_collection) + was_released = client.get_load_state(collection)['state'] == LoadState.NotLoad + try: + client.load_collection(collection, timeout=BM25_BACKFILL_LOAD_TIMEOUT) + _copy_to_bm25_collection(client, collection, staging_collection, output_fields, to_bm25_rows) + finally: + if was_released: + client.release_collection(collection) + except Exception as e: + log.error('Error building BM25 collection for %s: %s', collection, e) + client.drop_collection(staging_collection) + + +def _copy_to_bm25_collection( + client: Client, + collection: str, + staging_collection: str, + output_fields: list[str], + to_bm25_rows: Callable[[list[dict]], list[dict]], +): + last_id = None + for batch_size in BM25_BACKFILL_BATCH_SIZES: + try: + # Rows come in primary key order, so a failed read resumes after the last copied id. + iterator = client.query_iterator( + collection_name=collection, + filter='' if last_id is None else f"id > '{_escape_milvus_string(last_id)}'", + output_fields=output_fields, + batch_size=batch_size, + ) + while batch := iterator.next(): + for rows in _size_capped_batches(to_bm25_rows(batch)): + client.insert(collection_name=staging_collection, data=rows) + last_id = rows[-1]['id'] + iterator.close() + client.rename_collection(staging_collection, f'{collection}{BM25_COLLECTION_SUFFIX}') + return + except Exception as e: + log.warning('Error building BM25 collection for %s with batch size %s: %s', collection, batch_size, e) + log.error('Could not build BM25 collection for %s at any batch size.', collection) + client.drop_collection(staging_collection) + + class MilvusClient(VectorDBBase): def __init__(self): self.collection_prefix = 'open_webui' @@ -75,6 +178,8 @@ class MilvusClient(VectorDBBase): self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB) else: self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB, token=MILVUS_TOKEN) + if ENABLE_DB_MIGRATIONS: + self._backfill_bm25_collections() def _result_to_get_result(self, result) -> GetResult: ids = [] @@ -204,6 +309,60 @@ class MilvusClient(VectorDBBase): index_type, metric_type, ) + try: + self._create_bm25_collection(f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}') + except MilvusException as e: + log.warning( + 'Could not create BM25 collection for %s_%s (needs Milvus 2.5+): %s', + self.collection_prefix, + collection_name, + e, + ) + + def _create_bm25_collection(self, bm25_collection: str): + schema = self.client.create_schema(auto_id=False) + schema.add_field(field_name='id', datatype=DataType.VARCHAR, is_primary=True, max_length=65535) + schema.add_field( + field_name='text', datatype=DataType.VARCHAR, max_length=MILVUS_TEXT_MAX_LENGTH, enable_analyzer=True + ) + schema.add_field(field_name='sparse', datatype=DataType.SPARSE_FLOAT_VECTOR) + schema.add_field(field_name='metadata', datatype=DataType.JSON) + schema.add_function( + Function( + name='text_bm25', + function_type=FunctionType.BM25, + input_field_names=['text'], + output_field_names=['sparse'], + ) + ) + index_params = self.client.prepare_index_params() + index_params.add_index(field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25') + self.client.create_collection(collection_name=bm25_collection, schema=schema) + try: + self.client.create_index(collection_name=bm25_collection, index_params=index_params) + except MilvusException: + self.client.drop_collection(bm25_collection) + raise + + def _backfill_bm25_collections(self): + collection_names = set(self.client.list_collections()) + # Milvus Lite (a local .db file) breaks under concurrent collection changes. + max_workers = 1 if MILVUS_URI.endswith('.db') else BM25_BACKFILL_WORKERS + with ThreadPoolExecutor(max_workers=max_workers) as executor: + for collection_name_full in collection_names: + if ( + collection_name_full.startswith(f'{self.collection_prefix}_') + and not collection_name_full.endswith((BM25_COLLECTION_SUFFIX, BM25_STAGING_SUFFIX)) + and f'{collection_name_full}{BM25_COLLECTION_SUFFIX}' not in collection_names + ): + executor.submit( + _backfill_bm25_collection, + self.client, + collection_name_full, + self._create_bm25_collection, + ['id', 'data', 'metadata'], + _bm25_rows, + ) def has_collection(self, collection_name: str) -> bool: # Check if the collection exists based on the collection name. @@ -213,7 +372,11 @@ class MilvusClient(VectorDBBase): def delete_collection(self, collection_name: str): # Delete the collection based on the collection name. collection_name = collection_name.replace('-', '_') - return self.client.drop_collection(collection_name=f'{self.collection_prefix}_{collection_name}') + result = self.client.drop_collection(collection_name=f'{self.collection_prefix}_{collection_name}') + _update_bm25_collection( + self.client, f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}', 'drop_collection' + ) + return result def search( self, @@ -240,6 +403,62 @@ class MilvusClient(VectorDBBase): ) return self._result_to_search_result(result) + def hybrid_search( + self, + collection_name: str, + query: str, + vectors: list[list[float | int]], + filter: Optional[dict] = None, + limit: int = 10, + hybrid_bm25_weight: float = 0.5, + ) -> Optional[SearchResult]: + collection_name = collection_name.replace('-', '_') + bm25_collection = f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}' + if not self.client.has_collection(bm25_collection): + return None + self.client.load_collection(f'{self.collection_prefix}_{collection_name}') + + vector_result = None + if hybrid_bm25_weight < 1 and vectors: + vector_result = self.search(collection_name=collection_name, vectors=vectors, filter=filter, limit=limit) + + fts_results = [] + if hybrid_bm25_weight > 0 and query.strip(): + self.client.load_collection(bm25_collection) + metadata_exprs = _metadata_exprs(filter) + result = self.client.search( + collection_name=bm25_collection, + data=[query], + anns_field='sparse', + limit=limit, + filter=' and '.join(metadata_exprs), + ) + # The BM25 collection only ranks; text and metadata come from the vector collection. + id_list_str = ', '.join([f"'{_escape_milvus_string(str(hit['id']))}'" for hit in result[0]]) + items = self.client.query( + collection_name=f'{self.collection_prefix}_{collection_name}', + filter=' and '.join([*metadata_exprs, f'id in [{id_list_str}]']), + output_fields=['id', 'data', 'metadata'], + ) + items_by_id = {item['id']: item for item in items} + fts_results = [ + { + 'id': hit['id'], + 'text': (items_by_id[hit['id']].get('data') or {}).get('text'), + 'vmetadata': items_by_id[hit['id']]['metadata'], + } + for hit in result[0] + if hit['id'] in items_by_id + ] + + return merge_hybrid_search_results( + vector_result=vector_result, + fts_results=fts_results, + num_queries=len(vectors) or 1, + limit=limit, + hybrid_bm25_weight=hybrid_bm25_weight, + ) + def query(self, collection_name: str, filter: dict, limit: int = -1): collection_name = collection_name.replace('-', '_') if not self.has_collection(collection_name): @@ -332,10 +551,17 @@ class MilvusClient(VectorDBBase): } ) try: - return self.client.insert( + result = self.client.insert( collection_name=f'{self.collection_prefix}_{collection_name}', data=data, ) + _update_bm25_collection( + self.client, + f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}', + 'insert', + data=_bm25_rows(data), + ) + return result except MilvusException as e: log.error(f'Milvus insert failed for {self.collection_prefix}_{collection_name} ({len(items)} items): {e}') raise @@ -372,10 +598,17 @@ class MilvusClient(VectorDBBase): } ) try: - return self.client.upsert( + result = self.client.upsert( collection_name=f'{self.collection_prefix}_{collection_name}', data=data, ) + _update_bm25_collection( + self.client, + f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}', + 'upsert', + data=_bm25_rows(data), + ) + return result except MilvusException as e: log.error(f'Milvus upsert failed for {self.collection_prefix}_{collection_name} ({len(items)} items): {e}') raise @@ -392,12 +625,15 @@ class MilvusClient(VectorDBBase): log.warning(f'Delete attempted on non-existent collection: {self.collection_prefix}_{collection_name}') return None + bm25_collection = f'{self.collection_prefix}_{collection_name}{BM25_COLLECTION_SUFFIX}' if ids: log.info('Deleting items by IDs from %s_%s. IDs: %s', self.collection_prefix, collection_name, ids) - return self.client.delete( + result = self.client.delete( collection_name=f'{self.collection_prefix}_{collection_name}', ids=ids, ) + _update_bm25_collection(self.client, bm25_collection, 'delete', ids=ids) + return result elif filter: filter_string = ' && '.join( [f'metadata["{key}"] == {JSONCodec.dumps(value)}' for key, value in filter.items()] @@ -408,10 +644,12 @@ class MilvusClient(VectorDBBase): collection_name, filter_string, ) - return self.client.delete( + result = self.client.delete( collection_name=f'{self.collection_prefix}_{collection_name}', filter=filter_string, ) + _update_bm25_collection(self.client, bm25_collection, 'delete', filter=filter_string) + return result else: log.warning( f'Delete operation on {self.collection_prefix}_{collection_name} called without IDs or filter. No action taken.' diff --git a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py index de0a227ea3..1a998805f3 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py @@ -17,15 +17,22 @@ from open_webui.config import ( MILVUS_TOKEN, MILVUS_URI, ) -from open_webui.retrieval.vector.dbs.milvus import _metadata_exprs +from open_webui.env import ENABLE_DB_MIGRATIONS +from open_webui.retrieval.vector.dbs.milvus import ( + BM25_COLLECTION_SUFFIX, + BM25_STAGING_SUFFIX, + _backfill_bm25_collection, + _metadata_exprs, + _update_bm25_collection, +) from open_webui.retrieval.vector.main import ( GetResult, SearchResult, VectorDBBase, VectorItem, ) -from open_webui.retrieval.vector.utils import process_metadata -from pymilvus import DataType +from open_webui.retrieval.vector.utils import merge_hybrid_search_results, process_metadata +from pymilvus import DataType, Function, FunctionType from pymilvus import MilvusClient as Client from pymilvus.exceptions import MilvusException @@ -81,6 +88,8 @@ class MilvusClient(VectorDBBase): self.WEB_SEARCH_COLLECTION, self.HASH_BASED_COLLECTION, ] + if ENABLE_DB_MIGRATIONS: + self._backfill_bm25_collections() def _get_collection_and_resource_id(self, collection_name: str) -> Tuple[str, str]: """ @@ -132,6 +141,10 @@ class MilvusClient(VectorDBBase): self.client.create_collection(collection_name=mt_collection_name, schema=schema) self.client.create_index(collection_name=mt_collection_name, index_params=vector_index) + self._create_resource_id_index(mt_collection_name) + log.info('Created shared collection: %s', mt_collection_name) + + def _create_resource_id_index(self, mt_collection_name: str): try: # A Milvus server auto-selects the scalar index type from a parameterless call. self.client.create_index( @@ -148,11 +161,57 @@ class MilvusClient(VectorDBBase): # The index only accelerates resource_id filters; never fail # collection creation over it. log.warning(f'Could not create {RESOURCE_ID_FIELD} index on {mt_collection_name}: {e}') - log.info('Created shared collection: %s', mt_collection_name) + + def _create_bm25_collection(self, bm25_collection: str): + schema = self.client.create_schema(auto_id=False) + schema.add_field(field_name='id', datatype=DataType.VARCHAR, is_primary=True, max_length=36) + schema.add_field( + field_name='text', datatype=DataType.VARCHAR, max_length=MILVUS_TEXT_MAX_LENGTH, enable_analyzer=True + ) + schema.add_field(field_name='sparse', datatype=DataType.SPARSE_FLOAT_VECTOR) + schema.add_field(field_name='metadata', datatype=DataType.JSON) + schema.add_field(field_name=RESOURCE_ID_FIELD, datatype=DataType.VARCHAR, max_length=255) + schema.add_function( + Function( + name='text_bm25', + function_type=FunctionType.BM25, + input_field_names=['text'], + output_field_names=['sparse'], + ) + ) + self.client.create_collection(collection_name=bm25_collection, schema=schema) + try: + self.client.create_index( + collection_name=bm25_collection, + index_params=self.client.prepare_index_params( + field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25' + ), + ) + except MilvusException: + self.client.drop_collection(bm25_collection) + raise + self._create_resource_id_index(bm25_collection) + + def _backfill_bm25_collections(self): + for mt_collection in self.shared_collections: + if self.client.has_collection(mt_collection) and not self.client.has_collection( + f'{mt_collection}{BM25_COLLECTION_SUFFIX}' + ): + _backfill_bm25_collection( + self.client, + mt_collection, + self._create_bm25_collection, + ['id', 'text', 'metadata', RESOURCE_ID_FIELD], + lambda rows: rows, + ) def _ensure_collection(self, mt_collection_name: str, dimension: int): if not self.client.has_collection(mt_collection_name): self._create_shared_collection(mt_collection_name, dimension) + try: + self._create_bm25_collection(f'{mt_collection_name}{BM25_COLLECTION_SUFFIX}') + except MilvusException as e: + log.warning('Could not create BM25 collection for %s (needs Milvus 2.5+): %s', mt_collection_name, e) def has_collection(self, collection_name: str) -> bool: mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) @@ -199,6 +258,12 @@ class MilvusClient(VectorDBBase): try: self.client.insert(collection_name=mt_collection, data=entities) + _update_bm25_collection( + self.client, + f'{mt_collection}{BM25_COLLECTION_SUFFIX}', + 'insert', + data=[{key: value for key, value in entity.items() if key != 'vector'} for entity in entities], + ) except MilvusException as e: log.error( f'Milvus insert failed (collection={mt_collection}, ' @@ -250,6 +315,62 @@ class MilvusClient(VectorDBBase): return SearchResult(ids=ids, documents=documents, metadatas=metadatas, distances=distances) + def hybrid_search( + self, + collection_name: str, + query: str, + vectors: List[List[float]], + filter: Optional[Dict] = None, + limit: int = 10, + hybrid_bm25_weight: float = 0.5, + ) -> Optional[SearchResult]: + mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) + _validate_resource_id(resource_id) + bm25_collection = f'{mt_collection}{BM25_COLLECTION_SUFFIX}' + if not self.client.has_collection(bm25_collection): + return None + + vector_result = None + if hybrid_bm25_weight < 1 and vectors: + vector_result = self.search(collection_name=collection_name, vectors=vectors, filter=filter, limit=limit) + + fts_results = [] + if hybrid_bm25_weight > 0 and query.strip(): + self.client.load_collection(mt_collection) + self.client.load_collection(bm25_collection) + expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'", *_metadata_exprs(filter)] + results = self.client.search( + collection_name=bm25_collection, + data=[query], + anns_field='sparse', + limit=limit, + filter=' and '.join(expr), + ) + id_list_str = ', '.join([f"'{_escape_milvus_string(str(hit['id']))}'" for hit in results[0]]) + items = self.client.query( + collection_name=mt_collection, + filter=' and '.join([*expr, f'id in [{id_list_str}]']), + output_fields=['id', 'text', 'metadata'], + ) + items_by_id = {item['id']: item for item in items} + fts_results = [ + { + 'id': hit['id'], + 'text': items_by_id[hit['id']]['text'], + 'vmetadata': items_by_id[hit['id']]['metadata'], + } + for hit in results[0] + if hit['id'] in items_by_id + ] + + return merge_hybrid_search_results( + vector_result=vector_result, + fts_results=fts_results, + num_queries=len(vectors) or 1, + limit=limit, + hybrid_bm25_weight=hybrid_bm25_weight, + ) + def delete( self, collection_name: str, @@ -273,11 +394,16 @@ class MilvusClient(VectorDBBase): expr.append(f"metadata['{key}'] == '{_escape_milvus_string(str(value))}'") self.client.delete(collection_name=mt_collection, filter=' and '.join(expr)) + _update_bm25_collection( + self.client, f'{mt_collection}{BM25_COLLECTION_SUFFIX}', 'delete', filter=' and '.join(expr) + ) def reset(self): for collection_name in self.shared_collections: if self.client.has_collection(collection_name): self.client.drop_collection(collection_name) + self.client.drop_collection(f'{collection_name}{BM25_COLLECTION_SUFFIX}') + self.client.drop_collection(f'{collection_name}{BM25_STAGING_SUFFIX}') def delete_collection(self, collection_name: str): mt_collection, resource_id = self._get_collection_and_resource_id(collection_name) @@ -286,6 +412,12 @@ class MilvusClient(VectorDBBase): return self.client.delete(collection_name=mt_collection, filter=f"{RESOURCE_ID_FIELD} == '{resource_id}'") + _update_bm25_collection( + self.client, + f'{mt_collection}{BM25_COLLECTION_SUFFIX}', + 'delete', + filter=f"{RESOURCE_ID_FIELD} == '{resource_id}'", + ) def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]: mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)