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)