mirror of
https://github.com/open-webui/open-webui.git
synced 2026-08-28 05:27:35 +00:00
refac
Co-Authored-By: Classic298 <27028174+Classic298@users.noreply.github.com>
This commit is contained in:
parent
067114c280
commit
1d6d4e6e66
12 changed files with 191 additions and 35 deletions
|
|
@ -3,7 +3,7 @@ NOTE: This vector database integration is community-supported and maintained on
|
|||
"""
|
||||
|
||||
import ssl
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from elasticsearch import BadRequestError, Elasticsearch
|
||||
from elasticsearch.helpers import bulk, scan
|
||||
|
|
@ -23,7 +23,13 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
|
||||
|
||||
|
||||
def _metadata_filter(key: str, op: str, value: Any) -> dict:
|
||||
if op == '$in':
|
||||
return {'terms': {f'metadata.{key}': value}}
|
||||
return {'term': {f'metadata.{key}': value}}
|
||||
|
||||
|
||||
class ElasticsearchClient(VectorDBBase):
|
||||
|
|
@ -161,12 +167,16 @@ class ElasticsearchClient(VectorDBBase):
|
|||
filter: Optional[dict] = None,
|
||||
limit: int = 10,
|
||||
) -> Optional[SearchResult]:
|
||||
filters = [{'term': {'collection': collection_name}}]
|
||||
if filter:
|
||||
filters.extend(_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter))
|
||||
|
||||
query = {
|
||||
'size': limit,
|
||||
'_source': ['text', 'metadata'],
|
||||
'query': {
|
||||
'script_score': {
|
||||
'query': {'bool': {'filter': [{'term': {'collection': collection_name}}]}},
|
||||
'query': {'bool': {'filter': filters}},
|
||||
'script': {
|
||||
'source': "cosineSimilarity(params.vector, 'vector') + 1.0",
|
||||
'params': {'vector': vectors[0]}, # Assuming single query vector
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ NOTE: This vector database integration is community-supported and maintained on
|
|||
"""
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
import re
|
||||
from typing import Any, Optional
|
||||
|
||||
from open_webui.config import (
|
||||
MILVUS_DB,
|
||||
|
|
@ -23,7 +24,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
|
||||
from open_webui.utils.json_codec import JSONCodec
|
||||
from pymilvus import DataType
|
||||
from pymilvus import MilvusClient as Client
|
||||
|
|
@ -35,6 +36,36 @@ 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
|
||||
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
|
||||
|
||||
|
||||
def _escape_milvus_string(value: str) -> str:
|
||||
if not isinstance(value, str):
|
||||
raise TypeError(f'Expected str, got {type(value).__name__}')
|
||||
return value.replace('\\', '\\\\').replace("'", "\\'")
|
||||
|
||||
|
||||
def _milvus_literal(value: Any) -> str:
|
||||
if isinstance(value, str):
|
||||
return f"'{_escape_milvus_string(value)}'"
|
||||
if isinstance(value, bool):
|
||||
return str(value).lower()
|
||||
if isinstance(value, (int, float)):
|
||||
return str(value)
|
||||
raise TypeError(f'Unsupported Milvus filter value type: {type(value).__name__}')
|
||||
|
||||
|
||||
def _metadata_exprs(filter: Optional[dict]) -> list[str]:
|
||||
exprs = []
|
||||
for key, op, value in iter_filter_conditions(filter):
|
||||
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.fullmatch(key):
|
||||
raise ValueError(f'Invalid Milvus metadata filter key: {key!r}')
|
||||
if op == '$in':
|
||||
items = [f"metadata['{key}'] == {_milvus_literal(item)}" for item in value]
|
||||
exprs.append(f'({" or ".join(items)})' if items else 'false')
|
||||
else:
|
||||
exprs.append(f"metadata['{key}'] == {_milvus_literal(value)}")
|
||||
return exprs
|
||||
|
||||
|
||||
class MilvusClient(VectorDBBase):
|
||||
|
|
@ -193,6 +224,9 @@ class MilvusClient(VectorDBBase):
|
|||
) -> Optional[SearchResult]:
|
||||
# Search for the nearest neighbor items based on the vectors and return 'limit' number of results.
|
||||
collection_name = collection_name.replace('-', '_')
|
||||
kwargs = {}
|
||||
if filter:
|
||||
kwargs['filter'] = ' and '.join(_metadata_exprs(filter))
|
||||
# For some index types like IVF_FLAT, search params like nprobe can be set.
|
||||
# Example: search_params = {"nprobe": 10} if using IVF_FLAT
|
||||
# For simplicity, not adding configurable search_params here, but could be extended.
|
||||
|
|
@ -201,6 +235,7 @@ class MilvusClient(VectorDBBase):
|
|||
data=vectors,
|
||||
limit=limit,
|
||||
output_fields=['data', 'metadata'],
|
||||
**kwargs,
|
||||
# search_params=search_params # Potentially add later if needed
|
||||
)
|
||||
return self._result_to_search_result(result)
|
||||
|
|
@ -364,7 +399,9 @@ class MilvusClient(VectorDBBase):
|
|||
ids=ids,
|
||||
)
|
||||
elif filter:
|
||||
filter_string = ' && '.join([f'metadata["{key}"] == {JSONCodec.dumps(value)}' for key, value in filter.items()])
|
||||
filter_string = ' && '.join(
|
||||
[f'metadata["{key}"] == {JSONCodec.dumps(value)}' for key, value in filter.items()]
|
||||
)
|
||||
log.info(
|
||||
'Deleting items by filter from %s_%s. Filter: %s',
|
||||
self.collection_prefix,
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from open_webui.config import (
|
|||
MILVUS_TOKEN,
|
||||
MILVUS_URI,
|
||||
)
|
||||
from open_webui.retrieval.vector.dbs.milvus import _metadata_exprs
|
||||
from open_webui.retrieval.vector.main import (
|
||||
GetResult,
|
||||
SearchResult,
|
||||
|
|
@ -221,13 +222,14 @@ class MilvusClient(VectorDBBase):
|
|||
|
||||
self.client.load_collection(mt_collection)
|
||||
|
||||
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'", *_metadata_exprs(filter)]
|
||||
results = self.client.search(
|
||||
collection_name=mt_collection,
|
||||
data=vectors,
|
||||
anns_field='vector',
|
||||
search_params={'metric_type': MILVUS_METRIC_TYPE, 'params': {}},
|
||||
limit=limit,
|
||||
filter=f"{RESOURCE_ID_FIELD} == '{resource_id}'",
|
||||
filter=' and '.join(expr),
|
||||
output_fields=['id', 'text', 'metadata'],
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -68,7 +68,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
|
||||
|
||||
VECTOR_LENGTH = OPENGAUSS_INITIALIZE_MAX_VECTOR_LENGTH
|
||||
Base = declarative_base()
|
||||
|
|
@ -86,6 +86,12 @@ class DocumentChunk(Base):
|
|||
vmetadata = Column(MutableDict.as_mutable(JSONB), nullable=True)
|
||||
|
||||
|
||||
def _metadata_clause(key: str, op: str, value: Any):
|
||||
if op == '$in':
|
||||
return DocumentChunk.vmetadata[key].astext.in_([str(v) for v in value])
|
||||
return DocumentChunk.vmetadata[key].astext == str(value)
|
||||
|
||||
|
||||
class OpenGaussClient(VectorDBBase):
|
||||
def __init__(self) -> None:
|
||||
if not OPENGAUSS_DB_URL:
|
||||
|
|
@ -242,10 +248,15 @@ class OpenGaussClient(VectorDBBase):
|
|||
DocumentChunk.vmetadata,
|
||||
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label('distance'),
|
||||
]
|
||||
where_clauses = [DocumentChunk.collection_name == collection_name]
|
||||
if filter:
|
||||
where_clauses.extend(
|
||||
_metadata_clause(key, op, value) for key, op, value in iter_filter_conditions(filter)
|
||||
)
|
||||
|
||||
subq = (
|
||||
select(*result_fields)
|
||||
.where(DocumentChunk.collection_name == collection_name)
|
||||
.where(*where_clauses)
|
||||
.order_by(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector))
|
||||
)
|
||||
if limit is not None:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
NOTE: This vector database integration is community-supported and maintained on a best-effort basis.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from open_webui.config import (
|
||||
OPENSEARCH_CERT_VERIFY,
|
||||
|
|
@ -17,11 +17,17 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
|
||||
from opensearchpy import OpenSearch
|
||||
from opensearchpy.helpers import bulk
|
||||
|
||||
|
||||
def _metadata_filter(key: str, op: str, value: Any) -> dict:
|
||||
if op == '$in':
|
||||
return {'terms': {f'metadata.{key}.keyword': value}}
|
||||
return {'term': {f'metadata.{key}.keyword': value}}
|
||||
|
||||
|
||||
class OpenSearchClient(VectorDBBase):
|
||||
def __init__(self):
|
||||
self.index_prefix = 'open_webui'
|
||||
|
|
@ -121,6 +127,8 @@ class OpenSearchClient(VectorDBBase):
|
|||
filter: Optional[dict] = None,
|
||||
limit: int = 10,
|
||||
) -> Optional[SearchResult]:
|
||||
filter_clauses = [_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter)]
|
||||
|
||||
try:
|
||||
if not self.has_collection(collection_name):
|
||||
return None
|
||||
|
|
@ -130,7 +138,7 @@ class OpenSearchClient(VectorDBBase):
|
|||
'_source': ['text', 'metadata'],
|
||||
'query': {
|
||||
'script_score': {
|
||||
'query': {'match_all': {}},
|
||||
'query': {'bool': {'filter': filter_clauses}} if filter_clauses else {'match_all': {}},
|
||||
'script': {
|
||||
'source': '(cosineSimilarity(params.query_value, doc[params.field]) + 1.0) / 2.0',
|
||||
'params': {
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ import array
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from decimal import Decimal
|
||||
|
|
@ -56,9 +57,29 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions
|
||||
from open_webui.utils.json_codec import JSONCodec
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
_SAFE_METADATA_KEY_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]{0,63}$')
|
||||
|
||||
|
||||
def _metadata_where(filter: Optional[dict]) -> tuple[str, dict[str, Any]]:
|
||||
clause = ''
|
||||
params: dict[str, Any] = {}
|
||||
for i, (key, op, value) in enumerate(iter_filter_conditions(filter)):
|
||||
if not isinstance(key, str) or not _SAFE_METADATA_KEY_RE.fullmatch(key):
|
||||
raise ValueError(f'Invalid Oracle metadata filter key: {key!r}')
|
||||
json_value = f"JSON_VALUE(dc.vmetadata, '$.{key}' RETURNING VARCHAR2(4096))"
|
||||
if op == '$in':
|
||||
names = [f'value_{i}_{j}' for j, _ in enumerate(value)]
|
||||
clause += f' AND {json_value} IN ({", ".join(f":{name}" for name in names)})' if names else ' AND 1 = 0'
|
||||
params.update({name: str(item) for name, item in zip(names, value)})
|
||||
else:
|
||||
name = f'value_{i}'
|
||||
clause += f' AND {json_value} = :{name}'
|
||||
params[name] = str(value)
|
||||
return clause, params
|
||||
|
||||
|
||||
class Oracle23aiClient(VectorDBBase):
|
||||
|
|
@ -549,6 +570,7 @@ class Oracle23aiClient(VectorDBBase):
|
|||
return None
|
||||
|
||||
num_queries = len(vectors)
|
||||
filter_clause, filter_params = _metadata_where(filter)
|
||||
|
||||
ids = [[] for _ in range(num_queries)]
|
||||
distances = [[] for _ in range(num_queries)]
|
||||
|
|
@ -561,12 +583,12 @@ class Oracle23aiClient(VectorDBBase):
|
|||
vector_blob = self._vector_to_blob(vector)
|
||||
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT dc.id, dc.text,
|
||||
f"""
|
||||
SELECT dc.id, dc.text,
|
||||
JSON_SERIALIZE(dc.vmetadata RETURNING VARCHAR2(4096)) as vmetadata,
|
||||
VECTOR_DISTANCE(dc.vector, :query_vector, COSINE) as distance
|
||||
FROM document_chunk dc
|
||||
WHERE dc.collection_name = :collection_name
|
||||
WHERE dc.collection_name = :collection_name{filter_clause}
|
||||
ORDER BY VECTOR_DISTANCE(dc.vector, :query_vector, COSINE)
|
||||
FETCH APPROX FIRST :limit ROWS ONLY
|
||||
""",
|
||||
|
|
@ -574,6 +596,7 @@ class Oracle23aiClient(VectorDBBase):
|
|||
'query_vector': vector_blob,
|
||||
'collection_name': collection_name,
|
||||
'limit': limit,
|
||||
**filter_params,
|
||||
},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -35,7 +35,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import normalize_filter, process_metadata
|
||||
|
||||
NO_LIMIT = 10000 # Reasonable limit to avoid overwhelming the system
|
||||
BATCH_SIZE = 100 # Recommended batch size for Pinecone operations
|
||||
|
|
@ -372,13 +372,15 @@ class PineconeClient(VectorDBBase):
|
|||
try:
|
||||
# Search using the first vector (assuming this is the intended behavior)
|
||||
query_vector = vectors[0]
|
||||
pinecone_filter = normalize_filter(filter)
|
||||
pinecone_filter['collection_name'] = collection_name_with_prefix
|
||||
|
||||
# Perform the search
|
||||
query_response = self.index.query(
|
||||
vector=query_vector,
|
||||
top_k=limit,
|
||||
include_metadata=True,
|
||||
filter={'collection_name': collection_name_with_prefix},
|
||||
filter=pinecone_filter,
|
||||
)
|
||||
|
||||
matches = getattr(query_response, 'matches', []) or []
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ NOTE: This vector database integration is community-supported and maintained on
|
|||
"""
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from open_webui.config import (
|
||||
|
|
@ -22,6 +22,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions
|
||||
from qdrant_client import QdrantClient as Qclient
|
||||
from qdrant_client.http.models import PointStruct
|
||||
from qdrant_client.models import models
|
||||
|
|
@ -31,6 +32,11 @@ NO_LIMIT = 999999999
|
|||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _metadata_filter(key: str, op: str, value: Any) -> models.FieldCondition:
|
||||
match = models.MatchAny(any=value) if op == '$in' else models.MatchValue(value=value)
|
||||
return models.FieldCondition(key=f'metadata.{key}', match=match)
|
||||
|
||||
|
||||
class QdrantClient(VectorDBBase):
|
||||
def __init__(self):
|
||||
self.collection_prefix = QDRANT_COLLECTION_PREFIX
|
||||
|
|
@ -152,10 +158,13 @@ class QdrantClient(VectorDBBase):
|
|||
if limit is None:
|
||||
limit = NO_LIMIT # otherwise qdrant would set limit to 10!
|
||||
|
||||
conditions = [_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter)]
|
||||
query_filter = models.Filter(must=conditions) if conditions else None
|
||||
query_response = self.client.query_points(
|
||||
collection_name=f'{self.collection_prefix}_{collection_name}',
|
||||
query=vectors[0],
|
||||
limit=limit,
|
||||
query_filter=query_filter,
|
||||
)
|
||||
get_result = self._result_to_get_result(query_response.points)
|
||||
return SearchResult(
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions
|
||||
from qdrant_client import QdrantClient as Qclient
|
||||
from qdrant_client.http.exceptions import UnexpectedResponse
|
||||
from qdrant_client.http.models import PointStruct
|
||||
|
|
@ -39,8 +40,9 @@ def _tenant_filter(tenant_id: str) -> models.FieldCondition:
|
|||
return models.FieldCondition(key=TENANT_ID_FIELD, match=models.MatchValue(value=tenant_id))
|
||||
|
||||
|
||||
def _metadata_filter(key: str, value: Any) -> models.FieldCondition:
|
||||
return models.FieldCondition(key=f'metadata.{key}', match=models.MatchValue(value=value))
|
||||
def _metadata_filter(key: str, op: str, value: Any) -> models.FieldCondition:
|
||||
match = models.MatchAny(any=value) if op == '$in' else models.MatchValue(value=value)
|
||||
return models.FieldCondition(key=f'metadata.{key}', match=match)
|
||||
|
||||
|
||||
class QdrantClient(VectorDBBase):
|
||||
|
|
@ -234,7 +236,7 @@ class QdrantClient(VectorDBBase):
|
|||
# whose payload omits an id (e.g. memories), leaving orphaned vectors.
|
||||
must_conditions.append(models.HasIdCondition(has_id=ids))
|
||||
elif filter:
|
||||
must_conditions += [_metadata_filter(k, v) for k, v in filter.items()]
|
||||
must_conditions += [_metadata_filter(k, '$eq', v) for k, v in filter.items()]
|
||||
|
||||
return self.client.delete(
|
||||
collection_name=mt_collection,
|
||||
|
|
@ -258,12 +260,14 @@ class QdrantClient(VectorDBBase):
|
|||
log.debug("Collection %s doesn't exist, search returns None", mt_collection)
|
||||
return None
|
||||
|
||||
tenant_filter = _tenant_filter(tenant_id)
|
||||
conditions = [_tenant_filter(tenant_id)]
|
||||
if filter:
|
||||
conditions.extend(_metadata_filter(key, op, value) for key, op, value in iter_filter_conditions(filter))
|
||||
query_response = self.client.query_points(
|
||||
collection_name=mt_collection,
|
||||
query=vectors[0],
|
||||
limit=limit,
|
||||
query_filter=models.Filter(must=[tenant_filter]),
|
||||
query_filter=models.Filter(must=conditions),
|
||||
)
|
||||
get_result = self._result_to_get_result(query_response.points)
|
||||
return SearchResult(
|
||||
|
|
@ -286,7 +290,7 @@ class QdrantClient(VectorDBBase):
|
|||
if limit is None:
|
||||
limit = NO_LIMIT
|
||||
tenant_filter = _tenant_filter(tenant_id)
|
||||
field_conditions = [_metadata_filter(k, v) for k, v in filter.items()]
|
||||
field_conditions = [_metadata_filter(k, '$eq', v) for k, v in filter.items()]
|
||||
combined_filter = models.Filter(must=[tenant_filter, *field_conditions])
|
||||
points = self.client.scroll(
|
||||
collection_name=mt_collection,
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import metadata_matches_filter, normalize_filter, process_metadata
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -309,6 +309,7 @@ class S3VectorClient(VectorDBBase):
|
|||
|
||||
try:
|
||||
log.info("Searching collection '%s' with %s query vectors, limit=%s", collection_name, len(vectors), limit)
|
||||
vector_filter = normalize_filter(filter)
|
||||
|
||||
# Initialize result lists
|
||||
all_ids = []
|
||||
|
|
@ -323,15 +324,18 @@ class S3VectorClient(VectorDBBase):
|
|||
# Prepare the query vector in S3 Vector format
|
||||
query_vector_dict = {'float32': [float(x) for x in query_vector]}
|
||||
|
||||
# Call S3 Vector query API
|
||||
response = self.client.query_vectors(
|
||||
vectorBucketName=self.bucket_name,
|
||||
indexName=collection_name,
|
||||
topK=limit,
|
||||
queryVector=query_vector_dict,
|
||||
returnMetadata=True,
|
||||
returnDistance=True,
|
||||
)
|
||||
request_params = {
|
||||
'vectorBucketName': self.bucket_name,
|
||||
'indexName': collection_name,
|
||||
'topK': limit,
|
||||
'queryVector': query_vector_dict,
|
||||
'returnMetadata': True,
|
||||
'returnDistance': True,
|
||||
}
|
||||
if vector_filter:
|
||||
request_params['filter'] = vector_filter
|
||||
|
||||
response = self.client.query_vectors(**request_params)
|
||||
|
||||
# Process results for this query
|
||||
query_ids = []
|
||||
|
|
@ -346,6 +350,9 @@ class S3VectorClient(VectorDBBase):
|
|||
vector_metadata = vector.get('metadata', {})
|
||||
vector_distance = vector.get('distance', 0.0)
|
||||
|
||||
if vector_filter and not metadata_matches_filter(vector_metadata, vector_filter):
|
||||
continue
|
||||
|
||||
# Extract document text from metadata
|
||||
document_text = ''
|
||||
if isinstance(vector_metadata, dict):
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ from open_webui.retrieval.vector.main import (
|
|||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
from open_webui.retrieval.vector.utils import iter_filter_conditions, process_metadata
|
||||
|
||||
|
||||
def _convert_uuids_to_strings(obj: Any) -> Any:
|
||||
|
|
@ -54,6 +54,20 @@ def _convert_uuids_to_strings(obj: Any) -> Any:
|
|||
return obj
|
||||
|
||||
|
||||
def _metadata_filter(filter: Optional[dict]) -> Any:
|
||||
clauses = []
|
||||
for key, op, value in iter_filter_conditions(filter):
|
||||
if op == '$in':
|
||||
clauses.append(
|
||||
weaviate.classes.query.Filter.any_of(
|
||||
[weaviate.classes.query.Filter.by_property(name=key).equal(item) for item in value]
|
||||
)
|
||||
)
|
||||
else:
|
||||
clauses.append(weaviate.classes.query.Filter.by_property(name=key).equal(value))
|
||||
return weaviate.classes.query.Filter.all_of(clauses) if len(clauses) > 1 else (clauses[0] if clauses else None)
|
||||
|
||||
|
||||
class WeaviateClient(VectorDBBase):
|
||||
def __init__(self):
|
||||
self.url = WEAVIATE_HTTP_HOST
|
||||
|
|
@ -168,6 +182,7 @@ class WeaviateClient(VectorDBBase):
|
|||
return None
|
||||
|
||||
collection = self.client.collections.get(sane_collection_name)
|
||||
weaviate_filter = _metadata_filter(filter)
|
||||
|
||||
result_ids, result_documents, result_metadatas, result_distances = (
|
||||
[],
|
||||
|
|
@ -181,6 +196,7 @@ class WeaviateClient(VectorDBBase):
|
|||
response = collection.query.near_vector(
|
||||
near_vector=vector_embedding,
|
||||
limit=limit,
|
||||
filters=weaviate_filter,
|
||||
return_metadata=weaviate.classes.query.MetadataQuery(distance=True),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -61,6 +61,33 @@ def process_metadata(
|
|||
return result
|
||||
|
||||
|
||||
def iter_filter_conditions(filter: dict[str, Any] | None):
|
||||
for key, value in (filter or {}).items():
|
||||
if isinstance(value, dict):
|
||||
if set(value) != {'$in'}:
|
||||
raise ValueError(f"Unsupported metadata filter for '{key}': {value}")
|
||||
yield key, '$in', list(value['$in'])
|
||||
else:
|
||||
yield key, '$eq', value
|
||||
|
||||
|
||||
def normalize_filter(filter: dict[str, Any] | None) -> dict[str, Any]:
|
||||
return {key: {'$in': value} if op == '$in' else value for key, op, value in iter_filter_conditions(filter)}
|
||||
|
||||
|
||||
def metadata_matches_filter(metadata: dict[str, Any], filter: dict[str, Any] | None) -> bool:
|
||||
if not isinstance(metadata, dict):
|
||||
return False
|
||||
for key, op, value in iter_filter_conditions(filter):
|
||||
actual = metadata.get(key)
|
||||
if op == '$in':
|
||||
if actual not in value:
|
||||
return False
|
||||
elif actual != value:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def merge_hybrid_search_results(
|
||||
vector_result: SearchResult | None,
|
||||
fts_results: list[dict[str, Any]],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue