From 1d6d4e6e6647e1d403438ede7bd9ba20bc4cc8f6 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Tue, 25 Aug 2026 16:27:17 -0400 Subject: [PATCH] refac Co-Authored-By: Classic298 <27028174+Classic298@users.noreply.github.com> --- .../retrieval/vector/dbs/elasticsearch.py | 16 +++++-- .../open_webui/retrieval/vector/dbs/milvus.py | 43 +++++++++++++++++-- .../vector/dbs/milvus_multitenancy.py | 4 +- .../retrieval/vector/dbs/opengauss.py | 15 ++++++- .../retrieval/vector/dbs/opensearch.py | 14 ++++-- .../retrieval/vector/dbs/oracle23ai.py | 29 +++++++++++-- .../retrieval/vector/dbs/pinecone.py | 6 ++- .../open_webui/retrieval/vector/dbs/qdrant.py | 11 ++++- .../vector/dbs/qdrant_multitenancy.py | 16 ++++--- .../retrieval/vector/dbs/s3vector.py | 27 +++++++----- .../retrieval/vector/dbs/weaviate.py | 18 +++++++- backend/open_webui/retrieval/vector/utils.py | 27 ++++++++++++ 12 files changed, 191 insertions(+), 35 deletions(-) diff --git a/backend/open_webui/retrieval/vector/dbs/elasticsearch.py b/backend/open_webui/retrieval/vector/dbs/elasticsearch.py index da3e6e93c8..f8178bdbac 100644 --- a/backend/open_webui/retrieval/vector/dbs/elasticsearch.py +++ b/backend/open_webui/retrieval/vector/dbs/elasticsearch.py @@ -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 diff --git a/backend/open_webui/retrieval/vector/dbs/milvus.py b/backend/open_webui/retrieval/vector/dbs/milvus.py index 0aa392f683..d09f01c16d 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus.py @@ -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, diff --git a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py index 4e3aad49ce..24f1864fd4 100644 --- a/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py +++ b/backend/open_webui/retrieval/vector/dbs/milvus_multitenancy.py @@ -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'], ) diff --git a/backend/open_webui/retrieval/vector/dbs/opengauss.py b/backend/open_webui/retrieval/vector/dbs/opengauss.py index e11e944f3d..cfb7166250 100644 --- a/backend/open_webui/retrieval/vector/dbs/opengauss.py +++ b/backend/open_webui/retrieval/vector/dbs/opengauss.py @@ -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: diff --git a/backend/open_webui/retrieval/vector/dbs/opensearch.py b/backend/open_webui/retrieval/vector/dbs/opensearch.py index 798802cd54..48c538a10d 100644 --- a/backend/open_webui/retrieval/vector/dbs/opensearch.py +++ b/backend/open_webui/retrieval/vector/dbs/opensearch.py @@ -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': { diff --git a/backend/open_webui/retrieval/vector/dbs/oracle23ai.py b/backend/open_webui/retrieval/vector/dbs/oracle23ai.py index e72b7e326a..77408081d5 100644 --- a/backend/open_webui/retrieval/vector/dbs/oracle23ai.py +++ b/backend/open_webui/retrieval/vector/dbs/oracle23ai.py @@ -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, }, ) diff --git a/backend/open_webui/retrieval/vector/dbs/pinecone.py b/backend/open_webui/retrieval/vector/dbs/pinecone.py index 4d13977f3f..d1b7363839 100644 --- a/backend/open_webui/retrieval/vector/dbs/pinecone.py +++ b/backend/open_webui/retrieval/vector/dbs/pinecone.py @@ -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 [] diff --git a/backend/open_webui/retrieval/vector/dbs/qdrant.py b/backend/open_webui/retrieval/vector/dbs/qdrant.py index 8462badf1f..e683bb40c2 100644 --- a/backend/open_webui/retrieval/vector/dbs/qdrant.py +++ b/backend/open_webui/retrieval/vector/dbs/qdrant.py @@ -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( diff --git a/backend/open_webui/retrieval/vector/dbs/qdrant_multitenancy.py b/backend/open_webui/retrieval/vector/dbs/qdrant_multitenancy.py index 20fd8210d6..00941d6bb0 100644 --- a/backend/open_webui/retrieval/vector/dbs/qdrant_multitenancy.py +++ b/backend/open_webui/retrieval/vector/dbs/qdrant_multitenancy.py @@ -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, diff --git a/backend/open_webui/retrieval/vector/dbs/s3vector.py b/backend/open_webui/retrieval/vector/dbs/s3vector.py index 37effd3456..5a92234f78 100644 --- a/backend/open_webui/retrieval/vector/dbs/s3vector.py +++ b/backend/open_webui/retrieval/vector/dbs/s3vector.py @@ -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): diff --git a/backend/open_webui/retrieval/vector/dbs/weaviate.py b/backend/open_webui/retrieval/vector/dbs/weaviate.py index a896d8ed5e..6d398d4ced 100644 --- a/backend/open_webui/retrieval/vector/dbs/weaviate.py +++ b/backend/open_webui/retrieval/vector/dbs/weaviate.py @@ -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), ) diff --git a/backend/open_webui/retrieval/vector/utils.py b/backend/open_webui/retrieval/vector/utils.py index 404fd34386..19bba9b1ef 100644 --- a/backend/open_webui/retrieval/vector/utils.py +++ b/backend/open_webui/retrieval/vector/utils.py @@ -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]],