Enabling ENABLE_MILVUS_MULTITENANCY_MODE with the default MILVUS_URI (embedded Milvus Lite at DATA_DIR/vector_db/milvus.db) fails on the first embedding write: _create_shared_collection calls collection.create_index(RESOURCE_ID_FIELD) with no index params. A Milvus server auto-selects a scalar index type in that case, but Milvus Lite rejects the call with "create_index missing required 'index_type' parameter", so shared collection creation raises and every embedding write 500s (memory add, file upload, knowledge writes). Keep the parameterless call as the first attempt so behavior on Milvus servers is unchanged, fall back to an explicit INVERTED scalar index, and if that also fails log a warning and continue. The scalar index only accelerates resource_id filters; inserts and filtered queries work without it, so a missing index must not break collection creation. Verified against embedded Milvus Lite: shared collections now create (with the warning), and memory add, file upload and memory query succeed end to end. Against a Milvus server the first attempt is identical to the current code, so nothing changes where it works today.
348 lines
13 KiB
Python
348 lines
13 KiB
Python
"""
|
|
NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
|
|
"""
|
|
|
|
import logging
|
|
import re
|
|
from typing import Any, Dict, List, Optional, Tuple
|
|
|
|
from open_webui.config import (
|
|
MILVUS_COLLECTION_PREFIX,
|
|
MILVUS_DB,
|
|
MILVUS_HNSW_EFCONSTRUCTION,
|
|
MILVUS_HNSW_M,
|
|
MILVUS_INDEX_TYPE,
|
|
MILVUS_IVF_FLAT_NLIST,
|
|
MILVUS_METRIC_TYPE,
|
|
MILVUS_TOKEN,
|
|
MILVUS_URI,
|
|
)
|
|
from open_webui.retrieval.vector.main import (
|
|
GetResult,
|
|
SearchResult,
|
|
VectorDBBase,
|
|
VectorItem,
|
|
)
|
|
from pymilvus import (
|
|
Collection,
|
|
CollectionSchema,
|
|
DataType,
|
|
FieldSchema,
|
|
connections,
|
|
utility,
|
|
)
|
|
from pymilvus.exceptions import MilvusException
|
|
|
|
log = logging.getLogger(__name__)
|
|
|
|
RESOURCE_ID_FIELD = 'resource_id'
|
|
# Milvus VARCHAR hard cap for the `text` field (see _create_shared_collection).
|
|
# Chunks longer than this are truncated 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 expressions are SQL-like strings with no parameterized-query API;
|
|
# values get interpolated into single-quoted literals. Reject anything that
|
|
# can't be a legitimate Open WebUI collection name.
|
|
_SAFE_RESOURCE_ID_RE = re.compile(r'^[A-Za-z0-9_-]{1,255}$')
|
|
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
|
|
|
|
|
|
def _validate_resource_id(resource_id: str) -> str:
|
|
if not isinstance(resource_id, str) or not _SAFE_RESOURCE_ID_RE.match(resource_id):
|
|
raise ValueError(f'Invalid Milvus resource_id (collection name): {resource_id!r}')
|
|
return resource_id
|
|
|
|
|
|
def _validate_metadata_key(key: str) -> str:
|
|
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.match(key):
|
|
raise ValueError(f'Invalid Milvus metadata filter key: {key!r}')
|
|
return key
|
|
|
|
|
|
def _escape_milvus_string(value: str) -> str:
|
|
if not isinstance(value, str):
|
|
raise TypeError(f'Expected str for Milvus expression value, got {type(value).__name__}')
|
|
return value.replace('\\', '\\\\').replace("'", "\\'")
|
|
|
|
|
|
class MilvusClient(VectorDBBase):
|
|
def __init__(self):
|
|
# Milvus collection names can only contain numbers, letters, and underscores.
|
|
self.collection_prefix = MILVUS_COLLECTION_PREFIX.replace('-', '_')
|
|
connections.connect(
|
|
alias='default',
|
|
uri=MILVUS_URI,
|
|
token=MILVUS_TOKEN,
|
|
db_name=MILVUS_DB,
|
|
)
|
|
|
|
# Main collection types for multi-tenancy
|
|
self.MEMORY_COLLECTION = f'{self.collection_prefix}_memories'
|
|
self.KNOWLEDGE_COLLECTION = f'{self.collection_prefix}_knowledge'
|
|
self.FILE_COLLECTION = f'{self.collection_prefix}_files'
|
|
self.WEB_SEARCH_COLLECTION = f'{self.collection_prefix}_web_search'
|
|
self.HASH_BASED_COLLECTION = f'{self.collection_prefix}_hash_based'
|
|
self.shared_collections = [
|
|
self.MEMORY_COLLECTION,
|
|
self.KNOWLEDGE_COLLECTION,
|
|
self.FILE_COLLECTION,
|
|
self.WEB_SEARCH_COLLECTION,
|
|
self.HASH_BASED_COLLECTION,
|
|
]
|
|
|
|
def _get_collection_and_resource_id(self, collection_name: str) -> Tuple[str, str]:
|
|
"""
|
|
Maps the traditional collection name to multi-tenant collection and resource ID.
|
|
|
|
WARNING: This mapping relies on current Open WebUI naming conventions for
|
|
collection names. If Open WebUI changes how it generates collection names
|
|
(e.g., "user-memory-" prefix, "file-" prefix, web search patterns, or hash
|
|
formats), this mapping will break and route data to incorrect collections.
|
|
POTENTIALLY CAUSING HUGE DATA CORRUPTION, DATA CONSISTENCY ISSUES AND INCORRECT
|
|
DATA MAPPING INSIDE THE DATABASE.
|
|
"""
|
|
resource_id = collection_name
|
|
|
|
if collection_name.startswith('user-memory-'):
|
|
return self.MEMORY_COLLECTION, resource_id
|
|
elif collection_name.startswith('file-'):
|
|
return self.FILE_COLLECTION, resource_id
|
|
elif collection_name.startswith('web-search-'):
|
|
return self.WEB_SEARCH_COLLECTION, resource_id
|
|
elif len(collection_name) == 63 and all(c in '0123456789abcdef' for c in collection_name):
|
|
return self.HASH_BASED_COLLECTION, resource_id
|
|
else:
|
|
return self.KNOWLEDGE_COLLECTION, resource_id
|
|
|
|
def _create_shared_collection(self, mt_collection_name: str, dimension: int):
|
|
fields = [
|
|
FieldSchema(
|
|
name='id',
|
|
dtype=DataType.VARCHAR,
|
|
is_primary=True,
|
|
auto_id=False,
|
|
max_length=36,
|
|
),
|
|
FieldSchema(name='vector', dtype=DataType.FLOAT_VECTOR, dim=dimension),
|
|
FieldSchema(name='text', dtype=DataType.VARCHAR, max_length=65535),
|
|
FieldSchema(name='metadata', dtype=DataType.JSON),
|
|
FieldSchema(name=RESOURCE_ID_FIELD, dtype=DataType.VARCHAR, max_length=255),
|
|
]
|
|
schema = CollectionSchema(fields, 'Shared collection for multi-tenancy')
|
|
collection = Collection(mt_collection_name, schema)
|
|
|
|
index_params = {
|
|
'metric_type': MILVUS_METRIC_TYPE,
|
|
'index_type': MILVUS_INDEX_TYPE,
|
|
'params': {},
|
|
}
|
|
if MILVUS_INDEX_TYPE == 'HNSW':
|
|
index_params['params'] = {
|
|
'M': MILVUS_HNSW_M,
|
|
'efConstruction': MILVUS_HNSW_EFCONSTRUCTION,
|
|
}
|
|
elif MILVUS_INDEX_TYPE == 'IVF_FLAT':
|
|
index_params['params'] = {'nlist': MILVUS_IVF_FLAT_NLIST}
|
|
|
|
collection.create_index('vector', index_params)
|
|
try:
|
|
# A Milvus server auto-selects the scalar index type; embedded
|
|
# Milvus Lite requires an explicit one.
|
|
collection.create_index(RESOURCE_ID_FIELD)
|
|
except MilvusException:
|
|
try:
|
|
collection.create_index(RESOURCE_ID_FIELD, {'index_type': 'INVERTED'})
|
|
except MilvusException as e:
|
|
# 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(f'Created shared collection: {mt_collection_name}')
|
|
return collection
|
|
|
|
def _ensure_collection(self, mt_collection_name: str, dimension: int):
|
|
if not utility.has_collection(mt_collection_name):
|
|
self._create_shared_collection(mt_collection_name, dimension)
|
|
|
|
def has_collection(self, collection_name: str) -> bool:
|
|
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
|
|
_validate_resource_id(resource_id)
|
|
if not utility.has_collection(mt_collection):
|
|
return False
|
|
|
|
collection = Collection(mt_collection)
|
|
collection.load()
|
|
res = collection.query(expr=f"{RESOURCE_ID_FIELD} == '{resource_id}'", limit=1)
|
|
return len(res) > 0
|
|
|
|
def upsert(self, collection_name: str, items: List[VectorItem]):
|
|
if not items:
|
|
return
|
|
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
|
|
_validate_resource_id(resource_id)
|
|
dimension = len(items[0]['vector'])
|
|
self._ensure_collection(mt_collection, dimension)
|
|
collection = Collection(mt_collection)
|
|
|
|
entities = []
|
|
for item in items:
|
|
text = item['text'] or ''
|
|
if len(text) > 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})'
|
|
)
|
|
text = text[:MILVUS_TEXT_MAX_LENGTH]
|
|
entities.append(
|
|
{
|
|
'id': item['id'],
|
|
'vector': item['vector'],
|
|
'text': text,
|
|
'metadata': item['metadata'],
|
|
RESOURCE_ID_FIELD: resource_id,
|
|
}
|
|
)
|
|
|
|
try:
|
|
collection.insert(entities)
|
|
except MilvusException as e:
|
|
log.error(
|
|
f'Milvus insert failed (collection={mt_collection}, '
|
|
f'resource_id={resource_id}, items={len(entities)}): {e}'
|
|
)
|
|
raise
|
|
|
|
def search(
|
|
self,
|
|
collection_name: str,
|
|
vectors: List[List[float]],
|
|
filter: Optional[Dict] = None,
|
|
limit: int = 10,
|
|
) -> Optional[SearchResult]:
|
|
if not vectors:
|
|
return None
|
|
|
|
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
|
|
_validate_resource_id(resource_id)
|
|
if not utility.has_collection(mt_collection):
|
|
return None
|
|
|
|
collection = Collection(mt_collection)
|
|
collection.load()
|
|
|
|
search_params = {'metric_type': MILVUS_METRIC_TYPE, 'params': {}}
|
|
results = collection.search(
|
|
data=vectors,
|
|
anns_field='vector',
|
|
param=search_params,
|
|
limit=limit,
|
|
expr=f"{RESOURCE_ID_FIELD} == '{resource_id}'",
|
|
output_fields=['id', 'text', 'metadata'],
|
|
)
|
|
|
|
ids, documents, metadatas, distances = [], [], [], []
|
|
for hits in results:
|
|
batch_ids, batch_docs, batch_metadatas, batch_dists = [], [], [], []
|
|
for hit in hits:
|
|
batch_ids.append(hit.entity.get('id'))
|
|
batch_docs.append(hit.entity.get('text'))
|
|
batch_metadatas.append(hit.entity.get('metadata'))
|
|
batch_dists.append(hit.distance)
|
|
ids.append(batch_ids)
|
|
documents.append(batch_docs)
|
|
metadatas.append(batch_metadatas)
|
|
distances.append(batch_dists)
|
|
|
|
return SearchResult(ids=ids, documents=documents, metadatas=metadatas, distances=distances)
|
|
|
|
def delete(
|
|
self,
|
|
collection_name: str,
|
|
ids: Optional[List[str]] = None,
|
|
filter: Optional[Dict[str, Any]] = None,
|
|
):
|
|
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
|
|
_validate_resource_id(resource_id)
|
|
if not utility.has_collection(mt_collection):
|
|
return
|
|
|
|
collection = Collection(mt_collection)
|
|
|
|
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
|
|
if ids:
|
|
# Milvus expects a string list for 'in' operator
|
|
id_list_str = ', '.join([f"'{_escape_milvus_string(str(id_val))}'" for id_val in ids])
|
|
expr.append(f'id in [{id_list_str}]')
|
|
|
|
if filter:
|
|
for key, value in filter.items():
|
|
_validate_metadata_key(key)
|
|
expr.append(f"metadata['{key}'] == '{_escape_milvus_string(str(value))}'")
|
|
|
|
collection.delete(' and '.join(expr))
|
|
|
|
def reset(self):
|
|
for collection_name in self.shared_collections:
|
|
if utility.has_collection(collection_name):
|
|
utility.drop_collection(collection_name)
|
|
|
|
def delete_collection(self, collection_name: str):
|
|
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
|
|
_validate_resource_id(resource_id)
|
|
if not utility.has_collection(mt_collection):
|
|
return
|
|
|
|
collection = Collection(mt_collection)
|
|
collection.delete(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)
|
|
_validate_resource_id(resource_id)
|
|
if not utility.has_collection(mt_collection):
|
|
return None
|
|
|
|
collection = Collection(mt_collection)
|
|
collection.load()
|
|
|
|
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
|
|
if filter:
|
|
for key, value in filter.items():
|
|
_validate_metadata_key(key)
|
|
if isinstance(value, str):
|
|
expr.append(f"metadata['{key}'] == '{_escape_milvus_string(value)}'")
|
|
elif isinstance(value, bool):
|
|
expr.append(f"metadata['{key}'] == {str(value).lower()}")
|
|
elif isinstance(value, (int, float)):
|
|
expr.append(f"metadata['{key}'] == {value}")
|
|
else:
|
|
raise TypeError(f'Unsupported Milvus filter value type for key {key!r}: {type(value).__name__}')
|
|
|
|
iterator = collection.query_iterator(
|
|
expr=' and '.join(expr),
|
|
output_fields=['id', 'text', 'metadata'],
|
|
limit=limit if limit else -1,
|
|
)
|
|
|
|
all_results = []
|
|
while True:
|
|
batch = iterator.next()
|
|
if not batch:
|
|
iterator.close()
|
|
break
|
|
all_results.extend(batch)
|
|
|
|
ids = [res['id'] for res in all_results]
|
|
documents = [res['text'] for res in all_results]
|
|
metadatas = [res['metadata'] for res in all_results]
|
|
|
|
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
|
|
|
|
def get(self, collection_name: str) -> Optional[GetResult]:
|
|
return self.query(collection_name, filter={}, limit=None)
|
|
|
|
def insert(self, collection_name: str, items: List[VectorItem]):
|
|
return self.upsert(collection_name, items)
|