diff --git a/backend/open_webui/retrieval/vector/dbs/milvus.py b/backend/open_webui/retrieval/vector/dbs/milvus.py index d5aab4a1a8..0bf4a5955d 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus.py @@ -4,7 +4,7 @@ NOTE: This vector database integration is community-supported and maintained on import logging import re -from collections.abc import Callable, Iterator +from collections.abc import Callable, Iterable from concurrent.futures import ThreadPoolExecutor from typing import Any, Optional @@ -29,25 +29,23 @@ from open_webui.retrieval.vector.main import ( ) 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, Function, FunctionType +from pymilvus import CollectionSchema, DataType, Function, FunctionType from pymilvus import MilvusClient as Client from pymilvus.client.types import LoadState from pymilvus.exceptions import MilvusException log = logging.getLogger(__name__) -# Milvus caps stored text length (here the chunk lives under the JSON `data` +# Milvus caps stored text length (here the chunk lives in `text` or the JSON `data` # 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' +# Milvus cannot add BM25 to an existing collection, so migration copies each one into a new collection. 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 +# Even rows at Milvus's field size limits keep a batch this size below its 64 MB message limit. +BM25_BACKFILL_BATCH_SIZE = 128 BM25_BACKFILL_WORKERS = 8 -BM25_BACKFILL_LOAD_TIMEOUT = 300 +BM25_BACKFILL_INSERTS_IN_FLIGHT = 4 _SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$') @@ -80,95 +78,161 @@ def _metadata_exprs(filter: Optional[dict]) -> list[str]: return exprs +def _chunk_text(entity: dict) -> Optional[str]: + return entity['text'] if 'text' in entity else entity.get('data', {}).get('text') + + +def _truncate_text(text: str) -> str: + return text.encode()[:MILVUS_TEXT_MAX_LENGTH].decode(errors='ignore') + + 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 [ + { + 'id': row['id'], + 'vector': row['vector'], + 'text': row['data']['text'], + 'metadata': row['metadata'], + } + for row in rows + ] + + +def _add_bm25_fields(schema: CollectionSchema): + schema.add_field(field_name='sparse', datatype=DataType.SPARSE_FLOAT_VECTOR) + schema.add_function( + Function( + name='text_bm25', + function_type=FunctionType.BM25, + input_field_names=['text'], + output_field_names=['sparse'], ) - 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 _supports_bm25(client: Client) -> bool: + return tuple(int(part) for part in re.findall(r'\d+', client.get_server_version())[:2]) >= (2, 5) -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 _has_bm25_field(client: Client, collection: str) -> bool: + return any(field['name'] == 'sparse' for field in client.describe_collection(collection)['fields']) -def _backfill_bm25_collection( +def _pending_bm25_collections(client: Client, collections: Iterable[str], output_fields: list[str]) -> dict[str, int]: + """Finishes or clears an interrupted run, then maps each collection left to migrate to its dimension.""" + existing_collections = set(client.list_collections()) + pending_collections = {} + for collection in collections: + staging_collection = f'{collection}{BM25_STAGING_SUFFIX}' + if staging_collection in existing_collections: + if collection not in existing_collections: + # An earlier start dropped the original after a complete copy but stopped before this rename. + client.rename_collection(staging_collection, collection) + client.load_collection(collection) + continue + client.drop_collection(staging_collection) + if collection not in existing_collections: + continue + fields = client.describe_collection(collection)['fields'] + field_names = {field['name'] for field in fields} + if 'sparse' not in field_names and field_names.issuperset(output_fields): + pending_collections[collection] = next( + field['params']['dim'] for field in fields if field['name'] == 'vector' + ) + return pending_collections + + +def _backfill_bm25_collections( client: Client, - collection: str, - create_bm25_collection: Callable[[str], None], + collections: Iterable[str], + create_bm25_collection: Callable[[str, int], 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) + if not _supports_bm25(client): + log.info('Milvus has no BM25 (needs 2.5+), native hybrid search stays off.') + return + pending_collections = _pending_bm25_collections(client, collections, output_fields) + if not pending_collections: + return + + log.info('Migrating %s Milvus collections to native hybrid search.', len(pending_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: + copies = [ + executor.submit( + _copy_to_bm25_collection, + client, + collection, + dimension, + create_bm25_collection, + output_fields, + to_bm25_rows, + ) + for collection, dimension in pending_collections.items() + ] + copied = all(copy.result() for copy in copies) + if copied: + swaps = [ + executor.submit(_swap_in_bm25_collection, client, collection) for collection in pending_collections + ] + for swap in swaps: + swap.result() + else: + executor.shutdown(cancel_futures=True) + log.error('Milvus migration to native hybrid search failed, all collections are kept unchanged.') + for collection in pending_collections: + client.drop_collection(f'{collection}{BM25_STAGING_SUFFIX}') + return + log.info('Migrated %s Milvus collections to native hybrid search.', len(pending_collections)) + + +def _swap_in_bm25_collection(client: Client, collection: str): + client.drop_collection(collection) + client.rename_collection(f'{collection}{BM25_STAGING_SUFFIX}', collection) + client.load_collection(collection) def _copy_to_bm25_collection( client: Client, collection: str, - staging_collection: str, + dimension: int, + create_bm25_collection: Callable[[str, int], None], output_fields: list[str], to_bm25_rows: Callable[[list[dict]], list[dict]], -): - last_id = None - for batch_size in BM25_BACKFILL_BATCH_SIZES: +) -> bool: + staging_collection = f'{collection}{BM25_STAGING_SUFFIX}' + try: + create_bm25_collection(staging_collection, dimension) + was_released = client.get_load_state(collection)['state'] == LoadState.NotLoad try: - # Rows come in primary key order, so a failed read resumes after the last copied id. + client.load_collection(collection) 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, + batch_size=BM25_BACKFILL_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'] + with ThreadPoolExecutor(max_workers=BM25_BACKFILL_INSERTS_IN_FLIGHT) as insert_executor: + inserts = [] + while batch := iterator.next(): + if len(inserts) == BM25_BACKFILL_INSERTS_IN_FLIGHT: + inserts.pop(0).result() + inserts.append( + insert_executor.submit( + client.insert, collection_name=staging_collection, data=to_bm25_rows(batch) + ) + ) + for insert in inserts: + insert.result() 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) + finally: + if was_released: + client.release_collection(collection) + return True + except Exception as e: + log.error('Error copying Milvus collection %s: %s', collection, e) + return False class MilvusClient(VectorDBBase): @@ -179,7 +243,18 @@ class MilvusClient(VectorDBBase): else: self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB, token=MILVUS_TOKEN) if ENABLE_DB_MIGRATIONS: - self._backfill_bm25_collections() + collections = { + collection_name_full.removesuffix(BM25_STAGING_SUFFIX) + for collection_name_full in self.client.list_collections() + if collection_name_full.startswith(f'{self.collection_prefix}_') + } + _backfill_bm25_collections( + self.client, + collections, + self._create_unloaded_collection, + ['id', 'vector', 'data', 'metadata'], + _bm25_rows, + ) def _result_to_get_result(self, result) -> GetResult: ids = [] @@ -191,7 +266,7 @@ class MilvusClient(VectorDBBase): _metadatas = [] for item in match: _ids.append(item.get('id')) - _documents.append(item.get('data', {}).get('text')) + _documents.append(_chunk_text(item)) _metadatas.append(item.get('metadata')) ids.append(_ids) documents.append(_documents) @@ -220,7 +295,7 @@ class MilvusClient(VectorDBBase): # https://milvus.io/docs/de/metric.md _dist = (item.get('distance') + 1.0) / 2.0 _distances.append(_dist) - _documents.append(item.get('entity', {}).get('data', {}).get('text')) + _documents.append(_chunk_text(item.get('entity', {}))) _metadatas.append(item.get('entity', {}).get('metadata')) ids.append(_ids) distances.append(_distances) @@ -236,6 +311,12 @@ class MilvusClient(VectorDBBase): ) def _create_collection(self, collection_name: str, dimension: int): + collection_name_full = f'{self.collection_prefix}_{collection_name}' + self._create_unloaded_collection(collection_name_full, dimension) + self.client.load_collection(collection_name_full) + + def _create_unloaded_collection(self, collection_name_full: str, dimension: int): + supports_bm25 = _supports_bm25(self.client) schema = self.client.create_schema( auto_id=False, enable_dynamic_field=True, @@ -252,7 +333,17 @@ class MilvusClient(VectorDBBase): dim=dimension, description='vector', ) - schema.add_field(field_name='data', datatype=DataType.JSON, description='data') + if supports_bm25: + schema.add_field( + field_name='text', + datatype=DataType.VARCHAR, + max_length=MILVUS_TEXT_MAX_LENGTH, + enable_analyzer=True, + description='text', + ) + _add_bm25_fields(schema) + else: + schema.add_field(field_name='data', datatype=DataType.JSON, description='data') schema.add_field(field_name='metadata', datatype=DataType.JSON, description='metadata') index_params = self.client.prepare_index_params() @@ -297,72 +388,17 @@ class MilvusClient(VectorDBBase): params=index_creation_params, ) - self.client.create_collection( - collection_name=f'{self.collection_prefix}_{collection_name}', - schema=schema, - index_params=index_params, - ) + if supports_bm25: + index_params.add_index(field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25') + + self.client.create_collection(collection_name=collection_name_full, schema=schema) + self.client.create_index(collection_name=collection_name_full, index_params=index_params) log.info( - "Successfully created collection '%s_%s' with index type '%s' and metric '%s'.", - self.collection_prefix, - collection_name, + "Successfully created collection '%s' with index type '%s' and metric '%s'.", + collection_name_full, 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. @@ -372,11 +408,7 @@ class MilvusClient(VectorDBBase): def delete_collection(self, collection_name: str): # Delete the collection based on the collection name. collection_name = collection_name.replace('-', '_') - 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 + return self.client.drop_collection(collection_name=f'{self.collection_prefix}_{collection_name}') def search( self, @@ -396,8 +428,9 @@ class MilvusClient(VectorDBBase): result = self.client.search( collection_name=f'{self.collection_prefix}_{collection_name}', data=vectors, + anns_field='vector', limit=limit, - output_fields=['data', 'metadata'], + output_fields=['data', 'text', 'metadata'], **kwargs, # search_params=search_params # Potentially add later if needed ) @@ -413,8 +446,10 @@ class MilvusClient(VectorDBBase): 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): + collection_name_full = f'{self.collection_prefix}_{collection_name}' + if not self.client.has_collection(collection_name_full) or not _has_bm25_field( + self.client, collection_name_full + ): return None self.client.load_collection(f'{self.collection_prefix}_{collection_name}') @@ -424,31 +459,18 @@ class MilvusClient(VectorDBBase): 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, + collection_name=collection_name_full, data=[query], anns_field='sparse', limit=limit, filter=' and '.join(metadata_exprs), + output_fields=['text', 'metadata'], ) - # 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'], - } + {'id': hit['id'], 'text': hit['entity']['text'], 'vmetadata': hit['entity']['metadata']} for hit in result[0] - if hit['id'] in items_by_id ] return merge_hybrid_search_results( @@ -491,6 +513,7 @@ class MilvusClient(VectorDBBase): output_fields=[ 'id', 'data', + 'text', 'metadata', ], limit=limit if limit > 0 else -1, @@ -536,32 +559,31 @@ class MilvusClient(VectorDBBase): self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) log.info('Inserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) + has_bm25 = _has_bm25_field(self.client, f'{self.collection_prefix}_{collection_name}') data = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: - log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars') - text = text[:MILVUS_TEXT_MAX_LENGTH] - data.append( - { - 'id': item['id'], - 'vector': item['vector'], - 'data': {'text': text}, - 'metadata': process_metadata(item['metadata']), - } - ) + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: + log.warning( + 'Milvus: truncating text id=%s %s->%s bytes', item['id'], text_bytes, MILVUS_TEXT_MAX_LENGTH + ) + text = _truncate_text(text) + row = { + 'id': item['id'], + 'vector': item['vector'], + 'metadata': process_metadata(item['metadata']), + } + if has_bm25: + row['text'] = text + else: + row['data'] = {'text': text} + data.append(row) try: - result = self.client.insert( + return 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 @@ -583,32 +605,31 @@ class MilvusClient(VectorDBBase): self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector'])) log.info('Upserting %s items into collection %s_%s.', len(items), self.collection_prefix, collection_name) + has_bm25 = _has_bm25_field(self.client, f'{self.collection_prefix}_{collection_name}') data = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: - log.warning(f'Milvus: truncating text id={item["id"]} {len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars') - text = text[:MILVUS_TEXT_MAX_LENGTH] - data.append( - { - 'id': item['id'], - 'vector': item['vector'], - 'data': {'text': text}, - 'metadata': process_metadata(item['metadata']), - } - ) + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: + log.warning( + 'Milvus: truncating text id=%s %s->%s bytes', item['id'], text_bytes, MILVUS_TEXT_MAX_LENGTH + ) + text = _truncate_text(text) + row = { + 'id': item['id'], + 'vector': item['vector'], + 'metadata': process_metadata(item['metadata']), + } + if has_bm25: + row['text'] = text + else: + row['data'] = {'text': text} + data.append(row) try: - result = self.client.upsert( + return 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 @@ -625,15 +646,12 @@ 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) - result = self.client.delete( + return 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()] @@ -644,12 +662,10 @@ class MilvusClient(VectorDBBase): collection_name, filter_string, ) - result = self.client.delete( + return 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 1a998805f3..d0924fa9ae 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py @@ -19,11 +19,13 @@ from open_webui.config import ( ) 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, + _add_bm25_fields, + _backfill_bm25_collections, + _has_bm25_field, _metadata_exprs, - _update_bm25_collection, + _supports_bm25, + _truncate_text, ) from open_webui.retrieval.vector.main import ( GetResult, @@ -32,7 +34,7 @@ from open_webui.retrieval.vector.main import ( VectorItem, ) from open_webui.retrieval.vector.utils import merge_hybrid_search_results, process_metadata -from pymilvus import DataType, Function, FunctionType +from pymilvus import DataType from pymilvus import MilvusClient as Client from pymilvus.exceptions import MilvusException @@ -89,7 +91,13 @@ class MilvusClient(VectorDBBase): self.HASH_BASED_COLLECTION, ] if ENABLE_DB_MIGRATIONS: - self._backfill_bm25_collections() + _backfill_bm25_collections( + self.client, + self.shared_collections, + self._create_shared_collection, + ['id', 'vector', 'text', 'metadata', RESOURCE_ID_FIELD], + lambda rows: rows, + ) def _get_collection_and_resource_id(self, collection_name: str) -> Tuple[str, str]: """ @@ -116,10 +124,18 @@ class MilvusClient(VectorDBBase): return self.KNOWLEDGE_COLLECTION, resource_id def _create_shared_collection(self, mt_collection_name: str, dimension: int): + supports_bm25 = _supports_bm25(self.client) schema = self.client.create_schema(auto_id=False, description='Shared collection for multi-tenancy') schema.add_field(field_name='id', datatype=DataType.VARCHAR, is_primary=True, max_length=36) schema.add_field(field_name='vector', datatype=DataType.FLOAT_VECTOR, dim=dimension) - schema.add_field(field_name='text', datatype=DataType.VARCHAR, max_length=MILVUS_TEXT_MAX_LENGTH) + schema.add_field( + field_name='text', + datatype=DataType.VARCHAR, + max_length=MILVUS_TEXT_MAX_LENGTH, + enable_analyzer=supports_bm25, + ) + if supports_bm25: + _add_bm25_fields(schema) schema.add_field(field_name='metadata', datatype=DataType.JSON) schema.add_field(field_name=RESOURCE_ID_FIELD, datatype=DataType.VARCHAR, max_length=255) @@ -141,6 +157,13 @@ 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) + if supports_bm25: + self.client.create_index( + collection_name=mt_collection_name, + index_params=self.client.prepare_index_params( + field_name='sparse', index_type='SPARSE_INVERTED_INDEX', metric_type='BM25' + ), + ) self._create_resource_id_index(mt_collection_name) log.info('Created shared collection: %s', mt_collection_name) @@ -162,56 +185,9 @@ class MilvusClient(VectorDBBase): # collection creation over it. log.warning(f'Could not create {RESOURCE_ID_FIELD} index on {mt_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=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) @@ -239,13 +215,17 @@ class MilvusClient(VectorDBBase): entities = [] for item in items: text = item['text'] or '' - if len(text) > MILVUS_TEXT_MAX_LENGTH: + text_bytes = len(text.encode()) + if text_bytes > MILVUS_TEXT_MAX_LENGTH: log.warning( - f'Milvus: truncating text id={item["id"]} ' - f'{len(text)}->{MILVUS_TEXT_MAX_LENGTH} chars ' - f'(collection={mt_collection}, resource_id={resource_id})' + 'Milvus: truncating text id=%s %s->%s bytes (collection=%s, resource_id=%s)', + item['id'], + text_bytes, + MILVUS_TEXT_MAX_LENGTH, + mt_collection, + resource_id, ) - text = text[:MILVUS_TEXT_MAX_LENGTH] + text = _truncate_text(text) entities.append( { 'id': item['id'], @@ -258,12 +238,6 @@ 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}, ' @@ -326,8 +300,7 @@ class MilvusClient(VectorDBBase): ) -> 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): + if not self.client.has_collection(mt_collection) or not _has_bm25_field(self.client, mt_collection): return None vector_result = None @@ -337,30 +310,18 @@ class MilvusClient(VectorDBBase): 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, + collection_name=mt_collection, data=[query], anns_field='sparse', limit=limit, filter=' and '.join(expr), + output_fields=['text', 'metadata'], ) - 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'], - } + {'id': hit['id'], 'text': hit['entity']['text'], 'vmetadata': hit['entity']['metadata']} for hit in results[0] - if hit['id'] in items_by_id ] return merge_hybrid_search_results( @@ -394,15 +355,11 @@ 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): @@ -412,12 +369,6 @@ 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)