mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-05 02:41:34 +00:00
feat: Milvus hybrid search without a second text-only copy of every collection (#31660)
Native hybrid search on Milvus kept a second, text-only collection beside every Milvus collection and searched both. Each Milvus collection now holds its vectors, its text and its BM25 keyword index together, and results rank exactly as before, with the BM25 weight setting working the same way it does on pgvector. Existing data moves on the first start with ENABLE_DB_MIGRATIONS on: startup waits while every collection is copied once (vectors included, nothing is re-embedded) and the originals are only dropped after every copy succeeded, so a failed run changes nothing and is retried on the next start. The copy needs free disk space for a second copy of the data until it finishes, and on one 16-thread machine with Milvus's official docker compose setup it ran at 11 to 22 MB/s, so 200 GB takes about 2.5 to 5 hours depending on chunk size. Milvus servers older than 2.5 are detected and keep the existing hybrid search.
This commit is contained in:
parent
4ef7e35b88
commit
aee7c47277
2 changed files with 266 additions and 299 deletions
|
|
@ -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.'
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue