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
This commit is contained in:
Classic298 2026-09-30 19:03:47 +02:00 • committed by GitHub
parent 3d43a497b5
commit 75bff4bcd9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 381 additions and 11 deletions

View file

@ -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.'

View file

@ -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)