This commit is contained in:
Timothy Jaeryang Baek
2026-03-17 17:58:01 -05:00
parent fcf7208352
commit de3317e26b
220 changed files with 17200 additions and 22836 deletions
@@ -31,17 +31,15 @@ log = logging.getLogger(__name__)
class ChromaClient(VectorDBBase):
def __init__(self):
settings_dict = {
"allow_reset": True,
"anonymized_telemetry": False,
'allow_reset': True,
'anonymized_telemetry': False,
}
if CHROMA_CLIENT_AUTH_PROVIDER is not None:
settings_dict["chroma_client_auth_provider"] = CHROMA_CLIENT_AUTH_PROVIDER
settings_dict['chroma_client_auth_provider'] = CHROMA_CLIENT_AUTH_PROVIDER
if CHROMA_CLIENT_AUTH_CREDENTIALS is not None:
settings_dict["chroma_client_auth_credentials"] = (
CHROMA_CLIENT_AUTH_CREDENTIALS
)
settings_dict['chroma_client_auth_credentials'] = CHROMA_CLIENT_AUTH_CREDENTIALS
if CHROMA_HTTP_HOST != "":
if CHROMA_HTTP_HOST != '':
self.client = chromadb.HttpClient(
host=CHROMA_HTTP_HOST,
port=CHROMA_HTTP_PORT,
@@ -87,25 +85,23 @@ class ChromaClient(VectorDBBase):
# chromadb has cosine distance, 2 (worst) -> 0 (best). Re-odering to 0 -> 1
# https://docs.trychroma.com/docs/collections/configure cosine equation
distances: list = result["distances"][0]
distances: list = result['distances'][0]
distances = [2 - dist for dist in distances]
distances = [[dist / 2 for dist in distances]]
return SearchResult(
**{
"ids": result["ids"],
"distances": distances,
"documents": result["documents"],
"metadatas": result["metadatas"],
'ids': result['ids'],
'distances': distances,
'documents': result['documents'],
'metadatas': result['metadatas'],
}
)
return None
except Exception as e:
return None
def query(
self, collection_name: str, filter: dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]:
# Query the items from the collection based on the filter.
try:
collection = self.client.get_collection(name=collection_name)
@@ -117,9 +113,9 @@ class ChromaClient(VectorDBBase):
return GetResult(
**{
"ids": [result["ids"]],
"documents": [result["documents"]],
"metadatas": [result["metadatas"]],
'ids': [result['ids']],
'documents': [result['documents']],
'metadatas': [result['metadatas']],
}
)
return None
@@ -133,23 +129,21 @@ class ChromaClient(VectorDBBase):
result = collection.get()
return GetResult(
**{
"ids": [result["ids"]],
"documents": [result["documents"]],
"metadatas": [result["metadatas"]],
'ids': [result['ids']],
'documents': [result['documents']],
'metadatas': [result['metadatas']],
}
)
return None
def insert(self, collection_name: str, items: list[VectorItem]):
# Insert the items into the collection, if the collection does not exist, it will be created.
collection = self.client.get_or_create_collection(
name=collection_name, metadata={"hnsw:space": "cosine"}
)
collection = self.client.get_or_create_collection(name=collection_name, metadata={'hnsw:space': 'cosine'})
ids = [item["id"] for item in items]
documents = [item["text"] for item in items]
embeddings = [item["vector"] for item in items]
metadatas = [process_metadata(item["metadata"]) for item in items]
ids = [item['id'] for item in items]
documents = [item['text'] for item in items]
embeddings = [item['vector'] for item in items]
metadatas = [process_metadata(item['metadata']) for item in items]
for batch in create_batches(
api=self.client,
@@ -162,18 +156,14 @@ class ChromaClient(VectorDBBase):
def upsert(self, collection_name: str, items: list[VectorItem]):
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
collection = self.client.get_or_create_collection(
name=collection_name, metadata={"hnsw:space": "cosine"}
)
collection = self.client.get_or_create_collection(name=collection_name, metadata={'hnsw:space': 'cosine'})
ids = [item["id"] for item in items]
documents = [item["text"] for item in items]
embeddings = [item["vector"] for item in items]
metadatas = [process_metadata(item["metadata"]) for item in items]
ids = [item['id'] for item in items]
documents = [item['text'] for item in items]
embeddings = [item['vector'] for item in items]
metadatas = [process_metadata(item['metadata']) for item in items]
collection.upsert(
ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas
)
collection.upsert(ids=ids, documents=documents, embeddings=embeddings, metadatas=metadatas)
def delete(
self,
@@ -191,9 +181,7 @@ class ChromaClient(VectorDBBase):
collection.delete(where=filter)
except Exception as e:
# If collection doesn't exist, that's fine - nothing to delete
log.debug(
f"Attempted to delete from non-existent collection {collection_name}. Ignoring."
)
log.debug(f'Attempted to delete from non-existent collection {collection_name}. Ignoring.')
pass
def reset(self):
@@ -51,7 +51,7 @@ class ElasticsearchClient(VectorDBBase):
# Status: works
def _get_index_name(self, dimension: int) -> str:
return f"{self.index_prefix}_d{str(dimension)}"
return f'{self.index_prefix}_d{str(dimension)}'
# Status: works
def _scan_result_to_get_result(self, result) -> GetResult:
@@ -62,24 +62,24 @@ class ElasticsearchClient(VectorDBBase):
metadatas = []
for hit in result:
ids.append(hit["_id"])
documents.append(hit["_source"].get("text"))
metadatas.append(hit["_source"].get("metadata"))
ids.append(hit['_id'])
documents.append(hit['_source'].get('text'))
metadatas.append(hit['_source'].get('metadata'))
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
# Status: works
def _result_to_get_result(self, result) -> GetResult:
if not result["hits"]["hits"]:
if not result['hits']['hits']:
return None
ids = []
documents = []
metadatas = []
for hit in result["hits"]["hits"]:
ids.append(hit["_id"])
documents.append(hit["_source"].get("text"))
metadatas.append(hit["_source"].get("metadata"))
for hit in result['hits']['hits']:
ids.append(hit['_id'])
documents.append(hit['_source'].get('text'))
metadatas.append(hit['_source'].get('metadata'))
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
@@ -90,11 +90,11 @@ class ElasticsearchClient(VectorDBBase):
documents = []
metadatas = []
for hit in result["hits"]["hits"]:
ids.append(hit["_id"])
distances.append(hit["_score"])
documents.append(hit["_source"].get("text"))
metadatas.append(hit["_source"].get("metadata"))
for hit in result['hits']['hits']:
ids.append(hit['_id'])
distances.append(hit['_score'])
documents.append(hit['_source'].get('text'))
metadatas.append(hit['_source'].get('metadata'))
return SearchResult(
ids=[ids],
@@ -106,26 +106,26 @@ class ElasticsearchClient(VectorDBBase):
# Status: works
def _create_index(self, dimension: int):
body = {
"mappings": {
"dynamic_templates": [
'mappings': {
'dynamic_templates': [
{
"strings": {
"match_mapping_type": "string",
"mapping": {"type": "keyword"},
'strings': {
'match_mapping_type': 'string',
'mapping': {'type': 'keyword'},
}
}
],
"properties": {
"collection": {"type": "keyword"},
"id": {"type": "keyword"},
"vector": {
"type": "dense_vector",
"dims": dimension, # Adjust based on your vector dimensions
"index": True,
"similarity": "cosine",
'properties': {
'collection': {'type': 'keyword'},
'id': {'type': 'keyword'},
'vector': {
'type': 'dense_vector',
'dims': dimension, # Adjust based on your vector dimensions
'index': True,
'similarity': 'cosine',
},
"text": {"type": "text"},
"metadata": {"type": "object"},
'text': {'type': 'text'},
'metadata': {'type': 'object'},
},
}
}
@@ -139,21 +139,19 @@ class ElasticsearchClient(VectorDBBase):
# Status: works
def has_collection(self, collection_name) -> bool:
query_body = {"query": {"bool": {"filter": []}}}
query_body["query"]["bool"]["filter"].append(
{"term": {"collection": collection_name}}
)
query_body = {'query': {'bool': {'filter': []}}}
query_body['query']['bool']['filter'].append({'term': {'collection': collection_name}})
try:
result = self.client.count(index=f"{self.index_prefix}*", body=query_body)
result = self.client.count(index=f'{self.index_prefix}*', body=query_body)
return result.body["count"] > 0
return result.body['count'] > 0
except Exception as e:
return None
def delete_collection(self, collection_name: str):
query = {"query": {"term": {"collection": collection_name}}}
self.client.delete_by_query(index=f"{self.index_prefix}*", body=query)
query = {'query': {'term': {'collection': collection_name}}}
self.client.delete_by_query(index=f'{self.index_prefix}*', body=query)
# Status: works
def search(
@@ -164,51 +162,41 @@ class ElasticsearchClient(VectorDBBase):
limit: int = 10,
) -> Optional[SearchResult]:
query = {
"size": limit,
"_source": ["text", "metadata"],
"query": {
"script_score": {
"query": {
"bool": {"filter": [{"term": {"collection": collection_name}}]}
},
"script": {
"source": "cosineSimilarity(params.vector, 'vector') + 1.0",
"params": {
"vector": vectors[0]
}, # Assuming single query vector
'size': limit,
'_source': ['text', 'metadata'],
'query': {
'script_score': {
'query': {'bool': {'filter': [{'term': {'collection': collection_name}}]}},
'script': {
'source': "cosineSimilarity(params.vector, 'vector') + 1.0",
'params': {'vector': vectors[0]}, # Assuming single query vector
},
}
},
}
result = self.client.search(
index=self._get_index_name(len(vectors[0])), body=query
)
result = self.client.search(index=self._get_index_name(len(vectors[0])), body=query)
return self._result_to_search_result(result)
# Status: only tested halfwat
def query(
self, collection_name: str, filter: dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]:
if not self.has_collection(collection_name):
return None
query_body = {
"query": {"bool": {"filter": []}},
"_source": ["text", "metadata"],
'query': {'bool': {'filter': []}},
'_source': ['text', 'metadata'],
}
for field, value in filter.items():
query_body["query"]["bool"]["filter"].append({"term": {field: value}})
query_body["query"]["bool"]["filter"].append(
{"term": {"collection": collection_name}}
)
query_body['query']['bool']['filter'].append({'term': {field: value}})
query_body['query']['bool']['filter'].append({'term': {'collection': collection_name}})
size = limit if limit else 10
try:
result = self.client.search(
index=f"{self.index_prefix}*",
index=f'{self.index_prefix}*',
body=query_body,
size=size,
)
@@ -220,9 +208,7 @@ class ElasticsearchClient(VectorDBBase):
# Status: works
def _has_index(self, dimension: int):
return self.client.indices.exists(
index=self._get_index_name(dimension=dimension)
)
return self.client.indices.exists(index=self._get_index_name(dimension=dimension))
def get_or_create_index(self, dimension: int):
if not self._has_index(dimension=dimension):
@@ -232,28 +218,28 @@ class ElasticsearchClient(VectorDBBase):
def get(self, collection_name: str) -> Optional[GetResult]:
# Get all the items in the collection.
query = {
"query": {"bool": {"filter": [{"term": {"collection": collection_name}}]}},
"_source": ["text", "metadata"],
'query': {'bool': {'filter': [{'term': {'collection': collection_name}}]}},
'_source': ['text', 'metadata'],
}
results = list(scan(self.client, index=f"{self.index_prefix}*", query=query))
results = list(scan(self.client, index=f'{self.index_prefix}*', query=query))
return self._scan_result_to_get_result(results)
# Status: works
def insert(self, collection_name: str, items: list[VectorItem]):
if not self._has_index(dimension=len(items[0]["vector"])):
self._create_index(dimension=len(items[0]["vector"]))
if not self._has_index(dimension=len(items[0]['vector'])):
self._create_index(dimension=len(items[0]['vector']))
for batch in self._create_batches(items):
actions = [
{
"_index": self._get_index_name(dimension=len(items[0]["vector"])),
"_id": item["id"],
"_source": {
"collection": collection_name,
"vector": item["vector"],
"text": item["text"],
"metadata": process_metadata(item["metadata"]),
'_index': self._get_index_name(dimension=len(items[0]['vector'])),
'_id': item['id'],
'_source': {
'collection': collection_name,
'vector': item['vector'],
'text': item['text'],
'metadata': process_metadata(item['metadata']),
},
}
for item in batch
@@ -262,21 +248,21 @@ class ElasticsearchClient(VectorDBBase):
# Upsert documents using the update API with doc_as_upsert=True.
def upsert(self, collection_name: str, items: list[VectorItem]):
if not self._has_index(dimension=len(items[0]["vector"])):
self._create_index(dimension=len(items[0]["vector"]))
if not self._has_index(dimension=len(items[0]['vector'])):
self._create_index(dimension=len(items[0]['vector']))
for batch in self._create_batches(items):
actions = [
{
"_op_type": "update",
"_index": self._get_index_name(dimension=len(item["vector"])),
"_id": item["id"],
"doc": {
"collection": collection_name,
"vector": item["vector"],
"text": item["text"],
"metadata": process_metadata(item["metadata"]),
'_op_type': 'update',
'_index': self._get_index_name(dimension=len(item['vector'])),
'_id': item['id'],
'doc': {
'collection': collection_name,
'vector': item['vector'],
'text': item['text'],
'metadata': process_metadata(item['metadata']),
},
"doc_as_upsert": True,
'doc_as_upsert': True,
}
for item in batch
]
@@ -289,22 +275,17 @@ class ElasticsearchClient(VectorDBBase):
ids: Optional[list[str]] = None,
filter: Optional[dict] = None,
):
query = {
"query": {"bool": {"filter": [{"term": {"collection": collection_name}}]}}
}
query = {'query': {'bool': {'filter': [{'term': {'collection': collection_name}}]}}}
# logic based on chromaDB
if ids:
query["query"]["bool"]["filter"].append({"terms": {"_id": ids}})
query['query']['bool']['filter'].append({'terms': {'_id': ids}})
elif filter:
for field, value in filter.items():
query["query"]["bool"]["filter"].append(
{"term": {f"metadata.{field}": value}}
)
query['query']['bool']['filter'].append({'term': {f'metadata.{field}': value}})
self.client.delete_by_query(index=f"{self.index_prefix}*", body=query)
self.client.delete_by_query(index=f'{self.index_prefix}*', body=query)
def reset(self):
indices = self.client.indices.get(index=f"{self.index_prefix}*")
indices = self.client.indices.get(index=f'{self.index_prefix}*')
for index in indices:
self.client.indices.delete(index=index)
@@ -47,8 +47,8 @@ def _embedding_to_f32_bytes(vec: List[float]) -> bytes:
byte sequence. We use array('f') to avoid a numpy dependency and byteswap on
big-endian platforms for portability.
"""
a = array.array("f", [float(x) for x in vec]) # float32
if sys.byteorder != "little":
a = array.array('f', [float(x) for x in vec]) # float32
if sys.byteorder != 'little':
a.byteswap()
return a.tobytes()
@@ -68,7 +68,7 @@ def _safe_json(v: Any) -> Dict[str, Any]:
return v
if isinstance(v, (bytes, bytearray)):
try:
v = v.decode("utf-8")
v = v.decode('utf-8')
except Exception:
return {}
if isinstance(v, str):
@@ -105,16 +105,16 @@ class MariaDBVectorClient(VectorDBBase):
"""
self.db_url = (db_url or MARIADB_VECTOR_DB_URL).strip()
self.vector_length = int(vector_length)
self.distance_strategy = (distance_strategy or "cosine").strip().lower()
self.distance_strategy = (distance_strategy or 'cosine').strip().lower()
self.index_m = int(index_m)
if self.distance_strategy not in {"cosine", "euclidean"}:
if self.distance_strategy not in {'cosine', 'euclidean'}:
raise ValueError("distance_strategy must be 'cosine' or 'euclidean'")
if not self.db_url.lower().startswith("mariadb+mariadbconnector://"):
if not self.db_url.lower().startswith('mariadb+mariadbconnector://'):
raise ValueError(
"MariaDBVectorClient requires mariadb+mariadbconnector:// (official MariaDB driver) "
"to ensure qmark paramstyle and correct VECTOR binding."
'MariaDBVectorClient requires mariadb+mariadbconnector:// (official MariaDB driver) '
'to ensure qmark paramstyle and correct VECTOR binding.'
)
if isinstance(MARIADB_VECTOR_POOL_SIZE, int):
@@ -129,9 +129,7 @@ class MariaDBVectorClient(VectorDBBase):
poolclass=QueuePool,
)
else:
self.engine = create_engine(
self.db_url, pool_pre_ping=True, poolclass=NullPool
)
self.engine = create_engine(self.db_url, pool_pre_ping=True, poolclass=NullPool)
else:
self.engine = create_engine(self.db_url, pool_pre_ping=True)
self._init_schema()
@@ -185,7 +183,7 @@ class MariaDBVectorClient(VectorDBBase):
conn.commit()
except Exception as e:
conn.rollback()
log.exception(f"Error during database initialization: {e}")
log.exception(f'Error during database initialization: {e}')
raise
def _check_vector_length(self) -> None:
@@ -197,19 +195,19 @@ class MariaDBVectorClient(VectorDBBase):
"""
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute("SHOW CREATE TABLE document_chunk")
cur.execute('SHOW CREATE TABLE document_chunk')
row = cur.fetchone()
if not row or len(row) < 2:
return
ddl = row[1]
m = re.search(r"vector\\((\\d+)\\)", ddl, flags=re.IGNORECASE)
m = re.search(r'vector\\((\\d+)\\)', ddl, flags=re.IGNORECASE)
if not m:
return
existing = int(m.group(1))
if existing != int(self.vector_length):
raise Exception(
f"VECTOR_LENGTH {self.vector_length} does not match existing vector column dimension {existing}. "
"Cannot change vector size after initialization without migrating the data."
f'VECTOR_LENGTH {self.vector_length} does not match existing vector column dimension {existing}. '
'Cannot change vector size after initialization without migrating the data.'
)
def adjust_vector_length(self, vector: List[float]) -> List[float]:
@@ -227,11 +225,7 @@ class MariaDBVectorClient(VectorDBBase):
"""
Return the MariaDB Vector distance function name for the configured strategy.
"""
return (
"vec_distance_cosine"
if self.distance_strategy == "cosine"
else "vec_distance_euclidean"
)
return 'vec_distance_cosine' if self.distance_strategy == 'cosine' else 'vec_distance_euclidean'
def _score_from_dist(self, dist: float) -> float:
"""
@@ -240,7 +234,7 @@ class MariaDBVectorClient(VectorDBBase):
- cosine: score ~= 1 - cosine_distance, clamped to [0, 1]
- euclidean: score = 1 / (1 + dist)
"""
if self.distance_strategy == "cosine":
if self.distance_strategy == 'cosine':
score = 1.0 - dist
if score < 0.0:
score = 0.0
@@ -260,48 +254,48 @@ class MariaDBVectorClient(VectorDBBase):
- {"$or": [ ... ]}
"""
if not expr or not isinstance(expr, dict):
return "", []
return '', []
if "$and" in expr:
if '$and' in expr:
parts: List[str] = []
params: List[Any] = []
for e in expr.get("$and") or []:
for e in expr.get('$and') or []:
s, p = self._build_filter_sql_qmark(e)
if s:
parts.append(s)
params.extend(p)
return ("(" + " AND ".join(parts) + ")") if parts else "", params
return ('(' + ' AND '.join(parts) + ')') if parts else '', params
if "$or" in expr:
if '$or' in expr:
parts: List[str] = []
params: List[Any] = []
for e in expr.get("$or") or []:
for e in expr.get('$or') or []:
s, p = self._build_filter_sql_qmark(e)
if s:
parts.append(s)
params.extend(p)
return ("(" + " OR ".join(parts) + ")") if parts else "", params
return ('(' + ' OR '.join(parts) + ')') if parts else '', params
clauses: List[str] = []
params: List[Any] = []
for key, value in expr.items():
if key.startswith("$"):
if key.startswith('$'):
continue
json_expr = f"JSON_UNQUOTE(JSON_EXTRACT(vmetadata, '$.{key}'))"
if isinstance(value, dict) and "$in" in value:
vals = [str(v) for v in (value.get("$in") or [])]
if isinstance(value, dict) and '$in' in value:
vals = [str(v) for v in (value.get('$in') or [])]
if not vals:
clauses.append("0=1")
clauses.append('0=1')
continue
ors = []
for v in vals:
ors.append(f"{json_expr} = ?")
ors.append(f'{json_expr} = ?')
params.append(v)
clauses.append("(" + " OR ".join(ors) + ")")
clauses.append('(' + ' OR '.join(ors) + ')')
else:
clauses.append(f"{json_expr} = ?")
clauses.append(f'{json_expr} = ?')
params.append(str(value))
return ("(" + " AND ".join(clauses) + ")") if clauses else "", params
return ('(' + ' AND '.join(clauses) + ')') if clauses else '', params
def insert(self, collection_name: str, items: List[VectorItem]) -> None:
"""
@@ -322,15 +316,15 @@ class MariaDBVectorClient(VectorDBBase):
"""
params: List[Tuple[Any, ...]] = []
for item in items:
v = self.adjust_vector_length(item["vector"])
v = self.adjust_vector_length(item['vector'])
emb = _embedding_to_f32_bytes(v)
meta = process_metadata(item.get("metadata") or {})
meta = process_metadata(item.get('metadata') or {})
params.append(
(
item["id"],
item['id'],
emb,
collection_name,
item.get("text"),
item.get('text'),
json.dumps(meta),
)
)
@@ -338,7 +332,7 @@ class MariaDBVectorClient(VectorDBBase):
conn.commit()
except Exception as e:
conn.rollback()
log.exception(f"Error during insert: {e}")
log.exception(f'Error during insert: {e}')
raise
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
@@ -365,15 +359,15 @@ class MariaDBVectorClient(VectorDBBase):
"""
params: List[Tuple[Any, ...]] = []
for item in items:
v = self.adjust_vector_length(item["vector"])
v = self.adjust_vector_length(item['vector'])
emb = _embedding_to_f32_bytes(v)
meta = process_metadata(item.get("metadata") or {})
meta = process_metadata(item.get('metadata') or {})
params.append(
(
item["id"],
item['id'],
emb,
collection_name,
item.get("text"),
item.get('text'),
json.dumps(meta),
)
)
@@ -381,7 +375,7 @@ class MariaDBVectorClient(VectorDBBase):
conn.commit()
except Exception as e:
conn.rollback()
log.exception(f"Error during upsert: {e}")
log.exception(f'Error during upsert: {e}')
raise
def search(
@@ -415,10 +409,10 @@ class MariaDBVectorClient(VectorDBBase):
with self._connect() as conn:
with conn.cursor() as cur:
fsql, fparams = self._build_filter_sql_qmark(filter or {})
where = "collection_name = ?"
where = 'collection_name = ?'
base_params: List[Any] = [collection_name]
if fsql:
where = where + " AND " + fsql
where = where + ' AND ' + fsql
base_params.extend(fparams)
sql = f"""
@@ -460,26 +454,24 @@ class MariaDBVectorClient(VectorDBBase):
metadatas=metadatas,
)
except Exception as e:
log.exception(f"[MARIADB_VECTOR] search() failed: {e}")
log.exception(f'[MARIADB_VECTOR] search() failed: {e}')
return None
def query(
self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
"""
Retrieve documents by metadata filter (non-vector query).
"""
with self._connect() as conn:
with conn.cursor() as cur:
fsql, fparams = self._build_filter_sql_qmark(filter or {})
where = "collection_name = ?"
where = 'collection_name = ?'
params: List[Any] = [collection_name]
if fsql:
where = where + " AND " + fsql
where = where + ' AND ' + fsql
params.extend(fparams)
sql = f"SELECT id, text, vmetadata FROM document_chunk WHERE {where}"
sql = f'SELECT id, text, vmetadata FROM document_chunk WHERE {where}'
if limit is not None:
sql += " LIMIT ?"
sql += ' LIMIT ?'
params.append(int(limit))
cur.execute(sql, params)
rows = cur.fetchall()
@@ -490,18 +482,16 @@ class MariaDBVectorClient(VectorDBBase):
metadatas = [[_safe_json(r[2]) for r in rows]]
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
def get(
self, collection_name: str, limit: Optional[int] = None
) -> Optional[GetResult]:
def get(self, collection_name: str, limit: Optional[int] = None) -> Optional[GetResult]:
"""
Retrieve documents in a collection without filtering (optionally limited).
"""
with self._connect() as conn:
with conn.cursor() as cur:
sql = "SELECT id, text, vmetadata FROM document_chunk WHERE collection_name = ?"
sql = 'SELECT id, text, vmetadata FROM document_chunk WHERE collection_name = ?'
params: List[Any] = [collection_name]
if limit is not None:
sql += " LIMIT ?"
sql += ' LIMIT ?'
params.append(int(limit))
cur.execute(sql, params)
rows = cur.fetchall()
@@ -526,12 +516,12 @@ class MariaDBVectorClient(VectorDBBase):
with self._connect() as conn:
with conn.cursor() as cur:
try:
where = ["collection_name = ?"]
where = ['collection_name = ?']
params: List[Any] = [collection_name]
if ids:
ph = ", ".join(["?"] * len(ids))
where.append(f"id IN ({ph})")
ph = ', '.join(['?'] * len(ids))
where.append(f'id IN ({ph})')
params.extend(ids)
if filter:
@@ -540,12 +530,12 @@ class MariaDBVectorClient(VectorDBBase):
where.append(fsql)
params.extend(fparams)
sql = "DELETE FROM document_chunk WHERE " + " AND ".join(where)
sql = 'DELETE FROM document_chunk WHERE ' + ' AND '.join(where)
cur.execute(sql, params)
conn.commit()
except Exception as e:
conn.rollback()
log.exception(f"Error during delete: {e}")
log.exception(f'Error during delete: {e}')
raise
def reset(self) -> None:
@@ -555,11 +545,11 @@ class MariaDBVectorClient(VectorDBBase):
with self._connect() as conn:
with conn.cursor() as cur:
try:
cur.execute("TRUNCATE TABLE document_chunk")
cur.execute('TRUNCATE TABLE document_chunk')
conn.commit()
except Exception as e:
conn.rollback()
log.exception(f"Error during reset: {e}")
log.exception(f'Error during reset: {e}')
raise
def has_collection(self, collection_name: str) -> bool:
@@ -570,7 +560,7 @@ class MariaDBVectorClient(VectorDBBase):
with self._connect() as conn:
with conn.cursor() as cur:
cur.execute(
"SELECT 1 FROM document_chunk WHERE collection_name = ? LIMIT 1",
'SELECT 1 FROM document_chunk WHERE collection_name = ? LIMIT 1',
(collection_name,),
)
return cur.fetchone() is not None
@@ -590,4 +580,4 @@ class MariaDBVectorClient(VectorDBBase):
try:
self.engine.dispose()
except Exception as e:
log.exception(f"Error during dispose the underlying SQLAlchemy engine: {e}")
log.exception(f'Error during dispose the underlying SQLAlchemy engine: {e}')
+90 -127
View File
@@ -35,7 +35,7 @@ log = logging.getLogger(__name__)
class MilvusClient(VectorDBBase):
def __init__(self):
self.collection_prefix = "open_webui"
self.collection_prefix = 'open_webui'
if MILVUS_TOKEN is None:
self.client = Client(uri=MILVUS_URI, db_name=MILVUS_DB)
else:
@@ -50,17 +50,17 @@ class MilvusClient(VectorDBBase):
_documents = []
_metadatas = []
for item in match:
_ids.append(item.get("id"))
_documents.append(item.get("data", {}).get("text"))
_metadatas.append(item.get("metadata"))
_ids.append(item.get('id'))
_documents.append(item.get('data', {}).get('text'))
_metadatas.append(item.get('metadata'))
ids.append(_ids)
documents.append(_documents)
metadatas.append(_metadatas)
return GetResult(
**{
"ids": ids,
"documents": documents,
"metadatas": metadatas,
'ids': ids,
'documents': documents,
'metadatas': metadatas,
}
)
@@ -75,23 +75,23 @@ class MilvusClient(VectorDBBase):
_documents = []
_metadatas = []
for item in match:
_ids.append(item.get("id"))
_ids.append(item.get('id'))
# normalize milvus score from [-1, 1] to [0, 1] range
# https://milvus.io/docs/de/metric.md
_dist = (item.get("distance") + 1.0) / 2.0
_dist = (item.get('distance') + 1.0) / 2.0
_distances.append(_dist)
_documents.append(item.get("entity", {}).get("data", {}).get("text"))
_metadatas.append(item.get("entity", {}).get("metadata"))
_documents.append(item.get('entity', {}).get('data', {}).get('text'))
_metadatas.append(item.get('entity', {}).get('metadata'))
ids.append(_ids)
distances.append(_distances)
documents.append(_documents)
metadatas.append(_metadatas)
return SearchResult(
**{
"ids": ids,
"distances": distances,
"documents": documents,
"metadatas": metadatas,
'ids': ids,
'distances': distances,
'documents': documents,
'metadatas': metadatas,
}
)
@@ -101,21 +101,19 @@ class MilvusClient(VectorDBBase):
enable_dynamic_field=True,
)
schema.add_field(
field_name="id",
field_name='id',
datatype=DataType.VARCHAR,
is_primary=True,
max_length=65535,
)
schema.add_field(
field_name="vector",
field_name='vector',
datatype=DataType.FLOAT_VECTOR,
dim=dimension,
description="vector",
)
schema.add_field(field_name="data", datatype=DataType.JSON, description="data")
schema.add_field(
field_name="metadata", datatype=DataType.JSON, description="metadata"
description='vector',
)
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()
@@ -123,44 +121,44 @@ class MilvusClient(VectorDBBase):
index_type = MILVUS_INDEX_TYPE.upper()
metric_type = MILVUS_METRIC_TYPE.upper()
log.info(f"Using Milvus index type: {index_type}, metric type: {metric_type}")
log.info(f'Using Milvus index type: {index_type}, metric type: {metric_type}')
index_creation_params = {}
if index_type == "HNSW":
if index_type == 'HNSW':
index_creation_params = {
"M": MILVUS_HNSW_M,
"efConstruction": MILVUS_HNSW_EFCONSTRUCTION,
'M': MILVUS_HNSW_M,
'efConstruction': MILVUS_HNSW_EFCONSTRUCTION,
}
log.info(f"HNSW params: {index_creation_params}")
elif index_type == "IVF_FLAT":
index_creation_params = {"nlist": MILVUS_IVF_FLAT_NLIST}
log.info(f"IVF_FLAT params: {index_creation_params}")
elif index_type == "DISKANN":
log.info(f'HNSW params: {index_creation_params}')
elif index_type == 'IVF_FLAT':
index_creation_params = {'nlist': MILVUS_IVF_FLAT_NLIST}
log.info(f'IVF_FLAT params: {index_creation_params}')
elif index_type == 'DISKANN':
index_creation_params = {
"max_degree": MILVUS_DISKANN_MAX_DEGREE,
"search_list_size": MILVUS_DISKANN_SEARCH_LIST_SIZE,
'max_degree': MILVUS_DISKANN_MAX_DEGREE,
'search_list_size': MILVUS_DISKANN_SEARCH_LIST_SIZE,
}
log.info(f"DISKANN params: {index_creation_params}")
elif index_type in ["FLAT", "AUTOINDEX"]:
log.info(f"Using {index_type} index with no specific build-time params.")
log.info(f'DISKANN params: {index_creation_params}')
elif index_type in ['FLAT', 'AUTOINDEX']:
log.info(f'Using {index_type} index with no specific build-time params.')
else:
log.warning(
f"Unsupported MILVUS_INDEX_TYPE: '{index_type}'. "
f"Supported types: HNSW, IVF_FLAT, DISKANN, FLAT, AUTOINDEX. "
f"Milvus will use its default for the collection if this type is not directly supported for index creation."
f'Supported types: HNSW, IVF_FLAT, DISKANN, FLAT, AUTOINDEX. '
f'Milvus will use its default for the collection if this type is not directly supported for index creation.'
)
# For unsupported types, pass the type directly to Milvus; it might handle it or use a default.
# If Milvus errors out, the user needs to correct the MILVUS_INDEX_TYPE env var.
index_params.add_index(
field_name="vector",
field_name='vector',
index_type=index_type,
metric_type=metric_type,
params=index_creation_params,
)
self.client.create_collection(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
schema=schema,
index_params=index_params,
)
@@ -170,17 +168,13 @@ class MilvusClient(VectorDBBase):
def has_collection(self, collection_name: str) -> bool:
# Check if the collection exists based on the collection name.
collection_name = collection_name.replace("-", "_")
return self.client.has_collection(
collection_name=f"{self.collection_prefix}_{collection_name}"
)
collection_name = collection_name.replace('-', '_')
return self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}')
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}"
)
collection_name = collection_name.replace('-', '_')
return self.client.drop_collection(collection_name=f'{self.collection_prefix}_{collection_name}')
def search(
self,
@@ -190,15 +184,15 @@ class MilvusClient(VectorDBBase):
limit: int = 10,
) -> Optional[SearchResult]:
# Search for the nearest neighbor items based on the vectors and return 'limit' number of results.
collection_name = collection_name.replace("-", "_")
collection_name = collection_name.replace('-', '_')
# 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.
result = self.client.search(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
data=vectors,
limit=limit,
output_fields=["data", "metadata"],
output_fields=['data', 'metadata'],
# search_params=search_params # Potentially add later if needed
)
return self._result_to_search_result(result)
@@ -206,11 +200,9 @@ class MilvusClient(VectorDBBase):
def query(self, collection_name: str, filter: dict, limit: int = -1):
connections.connect(uri=MILVUS_URI, token=MILVUS_TOKEN, db_name=MILVUS_DB)
collection_name = collection_name.replace("-", "_")
collection_name = collection_name.replace('-', '_')
if not self.has_collection(collection_name):
log.warning(
f"Query attempted on non-existent collection: {self.collection_prefix}_{collection_name}"
)
log.warning(f'Query attempted on non-existent collection: {self.collection_prefix}_{collection_name}')
return None
filter_expressions = []
@@ -220,9 +212,9 @@ class MilvusClient(VectorDBBase):
else:
filter_expressions.append(f'metadata["{key}"] == {value}')
filter_string = " && ".join(filter_expressions)
filter_string = ' && '.join(filter_expressions)
collection = Collection(f"{self.collection_prefix}_{collection_name}")
collection = Collection(f'{self.collection_prefix}_{collection_name}')
collection.load()
try:
@@ -233,9 +225,9 @@ class MilvusClient(VectorDBBase):
iterator = collection.query_iterator(
expr=filter_string,
output_fields=[
"id",
"data",
"metadata",
'id',
'data',
'metadata',
],
limit=limit if limit > 0 else -1,
)
@@ -248,7 +240,7 @@ class MilvusClient(VectorDBBase):
break
all_results.extend(batch)
log.debug(f"Total results from query: {len(all_results)}")
log.debug(f'Total results from query: {len(all_results)}')
return self._result_to_get_result([all_results] if all_results else [[]])
except Exception as e:
@@ -259,7 +251,7 @@ class MilvusClient(VectorDBBase):
def get(self, collection_name: str) -> Optional[GetResult]:
# Get all the items in the collection. This can be very resource-intensive for large collections.
collection_name = collection_name.replace("-", "_")
collection_name = collection_name.replace('-', '_')
log.warning(
f"Fetching ALL items from collection '{self.collection_prefix}_{collection_name}'. This might be slow for large collections."
)
@@ -269,35 +261,25 @@ class MilvusClient(VectorDBBase):
def insert(self, collection_name: str, items: list[VectorItem]):
# Insert the items into the collection, if the collection does not exist, it will be created.
collection_name = collection_name.replace("-", "_")
if not self.client.has_collection(
collection_name=f"{self.collection_prefix}_{collection_name}"
):
log.info(
f"Collection {self.collection_prefix}_{collection_name} does not exist. Creating now."
)
collection_name = collection_name.replace('-', '_')
if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'):
log.info(f'Collection {self.collection_prefix}_{collection_name} does not exist. Creating now.')
if not items:
log.error(
f"Cannot create collection {self.collection_prefix}_{collection_name} without items to determine dimension."
f'Cannot create collection {self.collection_prefix}_{collection_name} without items to determine dimension.'
)
raise ValueError(
"Cannot create Milvus collection without items to determine vector dimension."
)
self._create_collection(
collection_name=collection_name, dimension=len(items[0]["vector"])
)
raise ValueError('Cannot create Milvus collection without items to determine vector dimension.')
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
log.info(
f"Inserting {len(items)} items into collection {self.collection_prefix}_{collection_name}."
)
log.info(f'Inserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
return self.client.insert(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
data=[
{
"id": item["id"],
"vector": item["vector"],
"data": {"text": item["text"]},
"metadata": process_metadata(item["metadata"]),
'id': item['id'],
'vector': item['vector'],
'data': {'text': item['text']},
'metadata': process_metadata(item['metadata']),
}
for item in items
],
@@ -305,35 +287,27 @@ class MilvusClient(VectorDBBase):
def upsert(self, collection_name: str, items: list[VectorItem]):
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
collection_name = collection_name.replace("-", "_")
if not self.client.has_collection(
collection_name=f"{self.collection_prefix}_{collection_name}"
):
log.info(
f"Collection {self.collection_prefix}_{collection_name} does not exist for upsert. Creating now."
)
collection_name = collection_name.replace('-', '_')
if not self.client.has_collection(collection_name=f'{self.collection_prefix}_{collection_name}'):
log.info(f'Collection {self.collection_prefix}_{collection_name} does not exist for upsert. Creating now.')
if not items:
log.error(
f"Cannot create collection {self.collection_prefix}_{collection_name} for upsert without items to determine dimension."
f'Cannot create collection {self.collection_prefix}_{collection_name} for upsert without items to determine dimension.'
)
raise ValueError(
"Cannot create Milvus collection for upsert without items to determine vector dimension."
'Cannot create Milvus collection for upsert without items to determine vector dimension.'
)
self._create_collection(
collection_name=collection_name, dimension=len(items[0]["vector"])
)
self._create_collection(collection_name=collection_name, dimension=len(items[0]['vector']))
log.info(
f"Upserting {len(items)} items into collection {self.collection_prefix}_{collection_name}."
)
log.info(f'Upserting {len(items)} items into collection {self.collection_prefix}_{collection_name}.')
return self.client.upsert(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
data=[
{
"id": item["id"],
"vector": item["vector"],
"data": {"text": item["text"]},
"metadata": process_metadata(item["metadata"]),
'id': item['id'],
'vector': item['vector'],
'data': {'text': item['text']},
'metadata': process_metadata(item['metadata']),
}
for item in items
],
@@ -346,46 +320,35 @@ class MilvusClient(VectorDBBase):
filter: Optional[dict] = None,
):
# Delete the items from the collection based on the ids or filter.
collection_name = collection_name.replace("-", "_")
collection_name = collection_name.replace('-', '_')
if not self.has_collection(collection_name):
log.warning(
f"Delete attempted on non-existent collection: {self.collection_prefix}_{collection_name}"
)
log.warning(f'Delete attempted on non-existent collection: {self.collection_prefix}_{collection_name}')
return None
if ids:
log.info(
f"Deleting items by IDs from {self.collection_prefix}_{collection_name}. IDs: {ids}"
)
log.info(f'Deleting items by IDs from {self.collection_prefix}_{collection_name}. IDs: {ids}')
return self.client.delete(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
ids=ids,
)
elif filter:
filter_string = " && ".join(
[
f'metadata["{key}"] == {json.dumps(value)}'
for key, value in filter.items()
]
)
filter_string = ' && '.join([f'metadata["{key}"] == {json.dumps(value)}' for key, value in filter.items()])
log.info(
f"Deleting items by filter from {self.collection_prefix}_{collection_name}. Filter: {filter_string}"
f'Deleting items by filter from {self.collection_prefix}_{collection_name}. Filter: {filter_string}'
)
return self.client.delete(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
filter=filter_string,
)
else:
log.warning(
f"Delete operation on {self.collection_prefix}_{collection_name} called without IDs or filter. No action taken."
f'Delete operation on {self.collection_prefix}_{collection_name} called without IDs or filter. No action taken.'
)
return None
def reset(self):
# Resets the database. This will delete all collections and item entries that match the prefix.
log.warning(
f"Resetting Milvus: Deleting all collections with prefix '{self.collection_prefix}'."
)
log.warning(f"Resetting Milvus: Deleting all collections with prefix '{self.collection_prefix}'.")
collection_names = self.client.list_collections()
deleted_collections = []
for collection_name_full in collection_names:
@@ -393,7 +356,7 @@ class MilvusClient(VectorDBBase):
try:
self.client.drop_collection(collection_name=collection_name_full)
deleted_collections.append(collection_name_full)
log.info(f"Deleted collection: {collection_name_full}")
log.info(f'Deleted collection: {collection_name_full}')
except Exception as e:
log.error(f"Error deleting collection {collection_name_full}: {e}")
log.info(f"Milvus reset complete. Deleted collections: {deleted_collections}")
log.error(f'Error deleting collection {collection_name_full}: {e}')
log.info(f'Milvus reset complete. Deleted collections: {deleted_collections}')
@@ -33,26 +33,26 @@ from pymilvus import (
log = logging.getLogger(__name__)
RESOURCE_ID_FIELD = "resource_id"
RESOURCE_ID_FIELD = 'resource_id'
class MilvusClient(VectorDBBase):
def __init__(self):
# Milvus collection names can only contain numbers, letters, and underscores.
self.collection_prefix = MILVUS_COLLECTION_PREFIX.replace("-", "_")
self.collection_prefix = MILVUS_COLLECTION_PREFIX.replace('-', '_')
connections.connect(
alias="default",
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.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,
@@ -74,15 +74,13 @@ class MilvusClient(VectorDBBase):
"""
resource_id = collection_name
if collection_name.startswith("user-memory-"):
if collection_name.startswith('user-memory-'):
return self.MEMORY_COLLECTION, resource_id
elif collection_name.startswith("file-"):
elif collection_name.startswith('file-'):
return self.FILE_COLLECTION, resource_id
elif collection_name.startswith("web-search-"):
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
):
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
@@ -90,36 +88,36 @@ class MilvusClient(VectorDBBase):
def _create_shared_collection(self, mt_collection_name: str, dimension: int):
fields = [
FieldSchema(
name="id",
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='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")
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": {},
'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,
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}
elif MILVUS_INDEX_TYPE == 'IVF_FLAT':
index_params['params'] = {'nlist': MILVUS_IVF_FLAT_NLIST}
collection.create_index("vector", index_params)
collection.create_index('vector', index_params)
collection.create_index(RESOURCE_ID_FIELD)
log.info(f"Created shared collection: {mt_collection_name}")
log.info(f'Created shared collection: {mt_collection_name}')
return collection
def _ensure_collection(self, mt_collection_name: str, dimension: int):
@@ -127,9 +125,7 @@ class MilvusClient(VectorDBBase):
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
)
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
if not utility.has_collection(mt_collection):
return False
@@ -141,19 +137,17 @@ class MilvusClient(VectorDBBase):
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
)
dimension = len(items[0]["vector"])
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
dimension = len(items[0]['vector'])
self._ensure_collection(mt_collection, dimension)
collection = Collection(mt_collection)
entities = [
{
"id": item["id"],
"vector": item["vector"],
"text": item["text"],
"metadata": item["metadata"],
'id': item['id'],
'vector': item['vector'],
'text': item['text'],
'metadata': item['metadata'],
RESOURCE_ID_FIELD: resource_id,
}
for item in items
@@ -170,41 +164,37 @@ class MilvusClient(VectorDBBase):
if not vectors:
return None
mt_collection, resource_id = self._get_collection_and_resource_id(
collection_name
)
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
if not utility.has_collection(mt_collection):
return None
collection = Collection(mt_collection)
collection.load()
search_params = {"metric_type": MILVUS_METRIC_TYPE, "params": {}}
search_params = {'metric_type': MILVUS_METRIC_TYPE, 'params': {}}
results = collection.search(
data=vectors,
anns_field="vector",
anns_field='vector',
param=search_params,
limit=limit,
expr=f"{RESOURCE_ID_FIELD} == '{resource_id}'",
output_fields=["id", "text", "metadata"],
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_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
)
return SearchResult(ids=ids, documents=documents, metadatas=metadatas, distances=distances)
def delete(
self,
@@ -212,9 +202,7 @@ class MilvusClient(VectorDBBase):
ids: Optional[List[str]] = None,
filter: Optional[Dict[str, Any]] = None,
):
mt_collection, resource_id = self._get_collection_and_resource_id(
collection_name
)
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
if not utility.has_collection(mt_collection):
return
@@ -224,14 +212,14 @@ class MilvusClient(VectorDBBase):
expr = [f"{RESOURCE_ID_FIELD} == '{resource_id}'"]
if ids:
# Milvus expects a string list for 'in' operator
id_list_str = ", ".join([f"'{id_val}'" for id_val in ids])
expr.append(f"id in [{id_list_str}]")
id_list_str = ', '.join([f"'{id_val}'" for id_val in ids])
expr.append(f'id in [{id_list_str}]')
if filter:
for key, value in filter.items():
expr.append(f"metadata['{key}'] == '{value}'")
collection.delete(" and ".join(expr))
collection.delete(' and '.join(expr))
def reset(self):
for collection_name in self.shared_collections:
@@ -239,21 +227,15 @@ class MilvusClient(VectorDBBase):
utility.drop_collection(collection_name)
def delete_collection(self, collection_name: str):
mt_collection, resource_id = self._get_collection_and_resource_id(
collection_name
)
mt_collection, resource_id = self._get_collection_and_resource_id(collection_name)
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
)
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)
if not utility.has_collection(mt_collection):
return None
@@ -269,8 +251,8 @@ class MilvusClient(VectorDBBase):
expr.append(f"metadata['{key}'] == {value}")
iterator = collection.query_iterator(
expr=" and ".join(expr),
output_fields=["id", "text", "metadata"],
expr=' and '.join(expr),
output_fields=['id', 'text', 'metadata'],
limit=limit if limit else -1,
)
@@ -282,9 +264,9 @@ class MilvusClient(VectorDBBase):
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]
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])
@@ -36,17 +36,15 @@ from sqlalchemy.dialects import registry
class OpenGaussDialect(PGDialect_psycopg2):
name = "opengauss"
name = 'opengauss'
def _get_server_version_info(self, connection):
try:
version = connection.exec_driver_sql("SELECT version()").scalar()
version = connection.exec_driver_sql('SELECT version()').scalar()
if not version:
return (9, 0, 0)
match = re.search(
r"openGauss\s+(\d+)\.(\d+)\.(\d+)(?:-\w+)?", version, re.IGNORECASE
)
match = re.search(r'openGauss\s+(\d+)\.(\d+)\.(\d+)(?:-\w+)?', version, re.IGNORECASE)
if match:
return (int(match.group(1)), int(match.group(2)), int(match.group(3)))
@@ -56,7 +54,7 @@ class OpenGaussDialect(PGDialect_psycopg2):
# Register dialect
registry.register("opengauss", __name__, "OpenGaussDialect")
registry.register('opengauss', __name__, 'OpenGaussDialect')
from open_webui.retrieval.vector.utils import process_metadata
from open_webui.retrieval.vector.main import (
@@ -80,11 +78,11 @@ VECTOR_LENGTH = OPENGAUSS_INITIALIZE_MAX_VECTOR_LENGTH
Base = declarative_base()
log = logging.getLogger(__name__)
log.setLevel(SRC_LOG_LEVELS["RAG"])
log.setLevel(SRC_LOG_LEVELS['RAG'])
class DocumentChunk(Base):
__tablename__ = "document_chunk"
__tablename__ = 'document_chunk'
id = Column(Text, primary_key=True)
vector = Column(Vector(dim=VECTOR_LENGTH), nullable=True)
@@ -100,26 +98,24 @@ class OpenGaussClient(VectorDBBase):
self.session = ScopedSession
else:
engine_kwargs = {"pool_pre_ping": True, "dialect": OpenGaussDialect()}
engine_kwargs = {'pool_pre_ping': True, 'dialect': OpenGaussDialect()}
if isinstance(OPENGAUSS_POOL_SIZE, int) and OPENGAUSS_POOL_SIZE > 0:
engine_kwargs.update(
{
"pool_size": OPENGAUSS_POOL_SIZE,
"max_overflow": OPENGAUSS_POOL_MAX_OVERFLOW,
"pool_timeout": OPENGAUSS_POOL_TIMEOUT,
"pool_recycle": OPENGAUSS_POOL_RECYCLE,
"poolclass": QueuePool,
'pool_size': OPENGAUSS_POOL_SIZE,
'max_overflow': OPENGAUSS_POOL_MAX_OVERFLOW,
'pool_timeout': OPENGAUSS_POOL_TIMEOUT,
'pool_recycle': OPENGAUSS_POOL_RECYCLE,
'poolclass': QueuePool,
}
)
else:
engine_kwargs["poolclass"] = NullPool
engine_kwargs['poolclass'] = NullPool
engine = create_engine(OPENGAUSS_DB_URL, **engine_kwargs)
SessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=engine, expire_on_commit=False
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
self.session = scoped_session(SessionLocal)
try:
@@ -128,47 +124,42 @@ class OpenGaussClient(VectorDBBase):
self.session.execute(
text(
"CREATE INDEX IF NOT EXISTS idx_document_chunk_vector "
"ON document_chunk USING ivfflat (vector vector_cosine_ops) WITH (lists = 100);"
'CREATE INDEX IF NOT EXISTS idx_document_chunk_vector '
'ON document_chunk USING ivfflat (vector vector_cosine_ops) WITH (lists = 100);'
)
)
self.session.execute(
text(
"CREATE INDEX IF NOT EXISTS idx_document_chunk_collection_name "
"ON document_chunk (collection_name);"
'CREATE INDEX IF NOT EXISTS idx_document_chunk_collection_name ON document_chunk (collection_name);'
)
)
self.session.commit()
log.info("OpenGauss vector database initialization completed.")
log.info('OpenGauss vector database initialization completed.')
except Exception as e:
self.session.rollback()
log.exception(f"OpenGauss Initialization failed.: {e}")
log.exception(f'OpenGauss Initialization failed.: {e}')
raise
def check_vector_length(self) -> None:
metadata = MetaData()
try:
document_chunk_table = Table(
"document_chunk", metadata, autoload_with=self.session.bind
)
document_chunk_table = Table('document_chunk', metadata, autoload_with=self.session.bind)
except NoSuchTableError:
return
if "vector" in document_chunk_table.columns:
vector_column = document_chunk_table.columns["vector"]
if 'vector' in document_chunk_table.columns:
vector_column = document_chunk_table.columns['vector']
vector_type = vector_column.type
if isinstance(vector_type, Vector):
db_vector_length = vector_type.dim
if db_vector_length != VECTOR_LENGTH:
raise Exception(
f"Vector dimension mismatch: configured {VECTOR_LENGTH} vs. {db_vector_length} in the database."
f'Vector dimension mismatch: configured {VECTOR_LENGTH} vs. {db_vector_length} in the database.'
)
else:
raise Exception("The 'vector' column type is not Vector.")
else:
raise Exception(
"The 'vector' column does not exist in the 'document_chunk' table."
)
raise Exception("The 'vector' column does not exist in the 'document_chunk' table.")
def adjust_vector_length(self, vector: List[float]) -> List[float]:
current_length = len(vector)
@@ -182,55 +173,47 @@ class OpenGaussClient(VectorDBBase):
try:
new_items = []
for item in items:
vector = self.adjust_vector_length(item["vector"])
vector = self.adjust_vector_length(item['vector'])
new_chunk = DocumentChunk(
id=item["id"],
id=item['id'],
vector=vector,
collection_name=collection_name,
text=item["text"],
vmetadata=process_metadata(item["metadata"]),
text=item['text'],
vmetadata=process_metadata(item['metadata']),
)
new_items.append(new_chunk)
self.session.bulk_save_objects(new_items)
self.session.commit()
log.info(
f"Inserting {len(new_items)} items into collection '{collection_name}'."
)
log.info(f"Inserting {len(new_items)} items into collection '{collection_name}'.")
except Exception as e:
self.session.rollback()
log.exception(f"Failed to insert data: {e}")
log.exception(f'Failed to insert data: {e}')
raise
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
try:
for item in items:
vector = self.adjust_vector_length(item["vector"])
existing = (
self.session.query(DocumentChunk)
.filter(DocumentChunk.id == item["id"])
.first()
)
vector = self.adjust_vector_length(item['vector'])
existing = self.session.query(DocumentChunk).filter(DocumentChunk.id == item['id']).first()
if existing:
existing.vector = vector
existing.text = item["text"]
existing.vmetadata = process_metadata(item["metadata"])
existing.text = item['text']
existing.vmetadata = process_metadata(item['metadata'])
existing.collection_name = collection_name
else:
new_chunk = DocumentChunk(
id=item["id"],
id=item['id'],
vector=vector,
collection_name=collection_name,
text=item["text"],
vmetadata=process_metadata(item["metadata"]),
text=item['text'],
vmetadata=process_metadata(item['metadata']),
)
self.session.add(new_chunk)
self.session.commit()
log.info(
f"Inserting/updating {len(items)} items in collection '{collection_name}'."
)
log.info(f"Inserting/updating {len(items)} items in collection '{collection_name}'.")
except Exception as e:
self.session.rollback()
log.exception(f"Failed to insert or update data.: {e}")
log.exception(f'Failed to insert or update data.: {e}')
raise
def search(
@@ -250,35 +233,29 @@ class OpenGaussClient(VectorDBBase):
def vector_expr(vector):
return cast(array(vector), Vector(VECTOR_LENGTH))
qid_col = column("qid", Integer)
q_vector_col = column("q_vector", Vector(VECTOR_LENGTH))
qid_col = column('qid', Integer)
q_vector_col = column('q_vector', Vector(VECTOR_LENGTH))
query_vectors = (
values(qid_col, q_vector_col)
.data(
[(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)]
)
.alias("query_vectors")
.data([(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)])
.alias('query_vectors')
)
result_fields = [
DocumentChunk.id,
DocumentChunk.text,
DocumentChunk.vmetadata,
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label(
"distance"
),
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label('distance'),
]
subq = (
select(*result_fields)
.where(DocumentChunk.collection_name == collection_name)
.order_by(
DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)
)
.order_by(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector))
)
if limit is not None:
subq = subq.limit(limit)
subq = subq.lateral("result")
subq = subq.lateral('result')
stmt = (
select(
@@ -309,21 +286,15 @@ class OpenGaussClient(VectorDBBase):
metadatas[qid].append(row.vmetadata)
self.session.rollback()
return SearchResult(
ids=ids, distances=distances, documents=documents, metadatas=metadatas
)
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
except Exception as e:
self.session.rollback()
log.exception(f"Vector search failed: {e}")
log.exception(f'Vector search failed: {e}')
return None
def query(
self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
try:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
for key, value in filter.items():
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
@@ -344,16 +315,12 @@ class OpenGaussClient(VectorDBBase):
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
except Exception as e:
self.session.rollback()
log.exception(f"Conditional query failed: {e}")
log.exception(f'Conditional query failed: {e}')
return None
def get(
self, collection_name: str, limit: Optional[int] = None
) -> Optional[GetResult]:
def get(self, collection_name: str, limit: Optional[int] = None) -> Optional[GetResult]:
try:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
if limit is not None:
query = query.limit(limit)
@@ -370,7 +337,7 @@ class OpenGaussClient(VectorDBBase):
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
except Exception as e:
self.session.rollback()
log.exception(f"Failed to retrieve data: {e}")
log.exception(f'Failed to retrieve data: {e}')
return None
def delete(
@@ -380,32 +347,28 @@ class OpenGaussClient(VectorDBBase):
filter: Optional[Dict[str, Any]] = None,
) -> None:
try:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
if ids:
query = query.filter(DocumentChunk.id.in_(ids))
if filter:
for key, value in filter.items():
query = query.filter(
DocumentChunk.vmetadata[key].astext == str(value)
)
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
deleted = query.delete(synchronize_session=False)
self.session.commit()
log.info(f"Deleted {deleted} items from collection '{collection_name}'")
except Exception as e:
self.session.rollback()
log.exception(f"Failed to delete data: {e}")
log.exception(f'Failed to delete data: {e}')
raise
def reset(self) -> None:
try:
deleted = self.session.query(DocumentChunk).delete()
self.session.commit()
log.info(f"Reset completed. Deleted {deleted} items")
log.info(f'Reset completed. Deleted {deleted} items')
except Exception as e:
self.session.rollback()
log.exception(f"Reset failed: {e}")
log.exception(f'Reset failed: {e}')
raise
def close(self) -> None:
@@ -414,16 +377,14 @@ class OpenGaussClient(VectorDBBase):
def has_collection(self, collection_name: str) -> bool:
try:
exists = (
self.session.query(DocumentChunk)
.filter(DocumentChunk.collection_name == collection_name)
.first()
self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name).first()
is not None
)
self.session.rollback()
return exists
except Exception as e:
self.session.rollback()
log.exception(f"Failed to check collection existence: {e}")
log.exception(f'Failed to check collection existence: {e}')
return False
def delete_collection(self, collection_name: str) -> None:
@@ -24,7 +24,7 @@ from open_webui.config import (
class OpenSearchClient(VectorDBBase):
def __init__(self):
self.index_prefix = "open_webui"
self.index_prefix = 'open_webui'
self.client = OpenSearch(
hosts=[OPENSEARCH_URI],
use_ssl=OPENSEARCH_SSL,
@@ -33,25 +33,25 @@ class OpenSearchClient(VectorDBBase):
)
def _get_index_name(self, collection_name: str) -> str:
return f"{self.index_prefix}_{collection_name}"
return f'{self.index_prefix}_{collection_name}'
def _result_to_get_result(self, result) -> GetResult:
if not result["hits"]["hits"]:
if not result['hits']['hits']:
return None
ids = []
documents = []
metadatas = []
for hit in result["hits"]["hits"]:
ids.append(hit["_id"])
documents.append(hit["_source"].get("text"))
metadatas.append(hit["_source"].get("metadata"))
for hit in result['hits']['hits']:
ids.append(hit['_id'])
documents.append(hit['_source'].get('text'))
metadatas.append(hit['_source'].get('metadata'))
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
def _result_to_search_result(self, result) -> SearchResult:
if not result["hits"]["hits"]:
if not result['hits']['hits']:
return None
ids = []
@@ -59,11 +59,11 @@ class OpenSearchClient(VectorDBBase):
documents = []
metadatas = []
for hit in result["hits"]["hits"]:
ids.append(hit["_id"])
distances.append(hit["_score"])
documents.append(hit["_source"].get("text"))
metadatas.append(hit["_source"].get("metadata"))
for hit in result['hits']['hits']:
ids.append(hit['_id'])
distances.append(hit['_score'])
documents.append(hit['_source'].get('text'))
metadatas.append(hit['_source'].get('metadata'))
return SearchResult(
ids=[ids],
@@ -74,33 +74,31 @@ class OpenSearchClient(VectorDBBase):
def _create_index(self, collection_name: str, dimension: int):
body = {
"settings": {"index": {"knn": True}},
"mappings": {
"properties": {
"id": {"type": "keyword"},
"vector": {
"type": "knn_vector",
"dimension": dimension, # Adjust based on your vector dimensions
"index": True,
"similarity": "faiss",
"method": {
"name": "hnsw",
"space_type": "innerproduct", # Use inner product to approximate cosine similarity
"engine": "faiss",
"parameters": {
"ef_construction": 128,
"m": 16,
'settings': {'index': {'knn': True}},
'mappings': {
'properties': {
'id': {'type': 'keyword'},
'vector': {
'type': 'knn_vector',
'dimension': dimension, # Adjust based on your vector dimensions
'index': True,
'similarity': 'faiss',
'method': {
'name': 'hnsw',
'space_type': 'innerproduct', # Use inner product to approximate cosine similarity
'engine': 'faiss',
'parameters': {
'ef_construction': 128,
'm': 16,
},
},
},
"text": {"type": "text"},
"metadata": {"type": "object"},
'text': {'type': 'text'},
'metadata': {'type': 'object'},
}
},
}
self.client.indices.create(
index=self._get_index_name(collection_name), body=body
)
self.client.indices.create(index=self._get_index_name(collection_name), body=body)
def _create_batches(self, items: list[VectorItem], batch_size=100):
for i in range(0, len(items), batch_size):
@@ -128,46 +126,40 @@ class OpenSearchClient(VectorDBBase):
return None
query = {
"size": limit,
"_source": ["text", "metadata"],
"query": {
"script_score": {
"query": {"match_all": {}},
"script": {
"source": "(cosineSimilarity(params.query_value, doc[params.field]) + 1.0) / 2.0",
"params": {
"field": "vector",
"query_value": vectors[0],
'size': limit,
'_source': ['text', 'metadata'],
'query': {
'script_score': {
'query': {'match_all': {}},
'script': {
'source': '(cosineSimilarity(params.query_value, doc[params.field]) + 1.0) / 2.0',
'params': {
'field': 'vector',
'query_value': vectors[0],
}, # Assuming single query vector
},
}
},
}
result = self.client.search(
index=self._get_index_name(collection_name), body=query
)
result = self.client.search(index=self._get_index_name(collection_name), body=query)
return self._result_to_search_result(result)
except Exception as e:
return None
def query(
self, collection_name: str, filter: dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: dict, limit: Optional[int] = None) -> Optional[GetResult]:
if not self.has_collection(collection_name):
return None
query_body = {
"query": {"bool": {"filter": []}},
"_source": ["text", "metadata"],
'query': {'bool': {'filter': []}},
'_source': ['text', 'metadata'],
}
for field, value in filter.items():
query_body["query"]["bool"]["filter"].append(
{"term": {"metadata." + str(field) + ".keyword": value}}
)
query_body['query']['bool']['filter'].append({'term': {'metadata.' + str(field) + '.keyword': value}})
size = limit if limit else 10000
@@ -188,28 +180,24 @@ class OpenSearchClient(VectorDBBase):
self._create_index(collection_name, dimension)
def get(self, collection_name: str) -> Optional[GetResult]:
query = {"query": {"match_all": {}}, "_source": ["text", "metadata"]}
query = {'query': {'match_all': {}}, '_source': ['text', 'metadata']}
result = self.client.search(
index=self._get_index_name(collection_name), body=query
)
result = self.client.search(index=self._get_index_name(collection_name), body=query)
return self._result_to_get_result(result)
def insert(self, collection_name: str, items: list[VectorItem]):
self._create_index_if_not_exists(
collection_name=collection_name, dimension=len(items[0]["vector"])
)
self._create_index_if_not_exists(collection_name=collection_name, dimension=len(items[0]['vector']))
for batch in self._create_batches(items):
actions = [
{
"_op_type": "index",
"_index": self._get_index_name(collection_name),
"_id": item["id"],
"_source": {
"vector": item["vector"],
"text": item["text"],
"metadata": process_metadata(item["metadata"]),
'_op_type': 'index',
'_index': self._get_index_name(collection_name),
'_id': item['id'],
'_source': {
'vector': item['vector'],
'text': item['text'],
'metadata': process_metadata(item['metadata']),
},
}
for item in batch
@@ -218,22 +206,20 @@ class OpenSearchClient(VectorDBBase):
self.client.indices.refresh(index=self._get_index_name(collection_name))
def upsert(self, collection_name: str, items: list[VectorItem]):
self._create_index_if_not_exists(
collection_name=collection_name, dimension=len(items[0]["vector"])
)
self._create_index_if_not_exists(collection_name=collection_name, dimension=len(items[0]['vector']))
for batch in self._create_batches(items):
actions = [
{
"_op_type": "update",
"_index": self._get_index_name(collection_name),
"_id": item["id"],
"doc": {
"vector": item["vector"],
"text": item["text"],
"metadata": process_metadata(item["metadata"]),
'_op_type': 'update',
'_index': self._get_index_name(collection_name),
'_id': item['id'],
'doc': {
'vector': item['vector'],
'text': item['text'],
'metadata': process_metadata(item['metadata']),
},
"doc_as_upsert": True,
'doc_as_upsert': True,
}
for item in batch
]
@@ -249,27 +235,23 @@ class OpenSearchClient(VectorDBBase):
if ids:
actions = [
{
"_op_type": "delete",
"_index": self._get_index_name(collection_name),
"_id": id,
'_op_type': 'delete',
'_index': self._get_index_name(collection_name),
'_id': id,
}
for id in ids
]
bulk(self.client, actions)
elif filter:
query_body = {
"query": {"bool": {"filter": []}},
'query': {'bool': {'filter': []}},
}
for field, value in filter.items():
query_body["query"]["bool"]["filter"].append(
{"term": {"metadata." + str(field) + ".keyword": value}}
)
self.client.delete_by_query(
index=self._get_index_name(collection_name), body=query_body
)
query_body['query']['bool']['filter'].append({'term': {'metadata.' + str(field) + '.keyword': value}})
self.client.delete_by_query(index=self._get_index_name(collection_name), body=query_body)
self.client.indices.refresh(index=self._get_index_name(collection_name))
def reset(self):
indices = self.client.indices.get(index=f"{self.index_prefix}_*")
indices = self.client.indices.get(index=f'{self.index_prefix}_*')
for index in indices:
self.client.indices.delete(index=index)
@@ -94,15 +94,15 @@ class Oracle23aiClient(VectorDBBase):
self._create_dbcs_pool()
dsn = ORACLE_DB_DSN
log.info(f"Creating Connection Pool [{ORACLE_DB_USER}:**@{dsn}]")
log.info(f'Creating Connection Pool [{ORACLE_DB_USER}:**@{dsn}]')
with self.get_connection() as connection:
log.info(f"Connection version: {connection.version}")
log.info(f'Connection version: {connection.version}')
self._initialize_database(connection)
log.info("Oracle Vector Search initialization complete.")
log.info('Oracle Vector Search initialization complete.')
except Exception as e:
log.exception(f"Error during Oracle Vector Search initialization: {e}")
log.exception(f'Error during Oracle Vector Search initialization: {e}')
raise
def _create_adb_pool(self) -> None:
@@ -122,7 +122,7 @@ class Oracle23aiClient(VectorDBBase):
wallet_location=ORACLE_WALLET_DIR,
wallet_password=ORACLE_WALLET_PASSWORD,
)
log.info("Created ADB connection pool with wallet authentication.")
log.info('Created ADB connection pool with wallet authentication.')
def _create_dbcs_pool(self) -> None:
"""
@@ -138,7 +138,7 @@ class Oracle23aiClient(VectorDBBase):
max=ORACLE_DB_POOL_MAX,
increment=ORACLE_DB_POOL_INCREMENT,
)
log.info("Created DB connection pool with basic authentication.")
log.info('Created DB connection pool with basic authentication.')
def get_connection(self):
"""
@@ -155,13 +155,11 @@ class Oracle23aiClient(VectorDBBase):
return connection
except oracledb.DatabaseError as e:
(error_obj,) = e.args
log.exception(
f"Connection attempt {attempt + 1} failed: {error_obj.message}"
)
log.exception(f'Connection attempt {attempt + 1} failed: {error_obj.message}')
if attempt < max_retries - 1:
wait_time = 2**attempt
log.info(f"Retrying in {wait_time} seconds...")
log.info(f'Retrying in {wait_time} seconds...')
time.sleep(wait_time)
else:
raise
@@ -177,30 +175,30 @@ class Oracle23aiClient(VectorDBBase):
def _monitor():
while True:
try:
log.info("[HealthCheck] Running periodic DB health check...")
log.info('[HealthCheck] Running periodic DB health check...')
self.ensure_connection()
log.info("[HealthCheck] Connection is healthy.")
log.info('[HealthCheck] Connection is healthy.')
except Exception as e:
log.exception(f"[HealthCheck] Connection health check failed: {e}")
log.exception(f'[HealthCheck] Connection health check failed: {e}')
time.sleep(interval_seconds)
thread = threading.Thread(target=_monitor, daemon=True)
thread.start()
log.info(f"Started DB health monitor every {interval_seconds} seconds.")
log.info(f'Started DB health monitor every {interval_seconds} seconds.')
def _reconnect_pool(self):
"""
Attempt to reinitialize the connection pool if it's been closed or broken.
"""
try:
log.info("Attempting to reinitialize the Oracle connection pool...")
log.info('Attempting to reinitialize the Oracle connection pool...')
# Close existing pool if it exists
if self.pool:
try:
self.pool.close()
except Exception as close_error:
log.warning(f"Error closing existing pool: {close_error}")
log.warning(f'Error closing existing pool: {close_error}')
# Re-create the appropriate connection pool based on DB type
if ORACLE_DB_USE_WALLET:
@@ -208,9 +206,9 @@ class Oracle23aiClient(VectorDBBase):
else: # DBCS
self._create_dbcs_pool()
log.info("Connection pool reinitialized.")
log.info('Connection pool reinitialized.')
except Exception as e:
log.exception(f"Failed to reinitialize the connection pool: {e}")
log.exception(f'Failed to reinitialize the connection pool: {e}')
raise
def ensure_connection(self):
@@ -220,11 +218,9 @@ class Oracle23aiClient(VectorDBBase):
try:
with self.get_connection() as connection:
with connection.cursor() as cursor:
cursor.execute("SELECT 1 FROM dual")
cursor.execute('SELECT 1 FROM dual')
except Exception as e:
log.exception(
f"Connection check failed: {e}, attempting to reconnect pool..."
)
log.exception(f'Connection check failed: {e}, attempting to reconnect pool...')
self._reconnect_pool()
def _output_type_handler(self, cursor, metadata):
@@ -239,9 +235,7 @@ class Oracle23aiClient(VectorDBBase):
A variable with appropriate conversion for vector types
"""
if metadata.type_code is oracledb.DB_TYPE_VECTOR:
return cursor.var(
metadata.type_code, arraysize=cursor.arraysize, outconverter=list
)
return cursor.var(metadata.type_code, arraysize=cursor.arraysize, outconverter=list)
def _initialize_database(self, connection) -> None:
"""
@@ -257,7 +251,7 @@ class Oracle23aiClient(VectorDBBase):
"""
with connection.cursor() as cursor:
try:
log.info("Creating Table document_chunk")
log.info('Creating Table document_chunk')
cursor.execute(
"""
BEGIN
@@ -279,7 +273,7 @@ class Oracle23aiClient(VectorDBBase):
"""
)
log.info("Creating Index document_chunk_collection_name_idx")
log.info('Creating Index document_chunk_collection_name_idx')
cursor.execute(
"""
BEGIN
@@ -296,7 +290,7 @@ class Oracle23aiClient(VectorDBBase):
"""
)
log.info("Creating VECTOR INDEX document_chunk_vector_ivf_idx")
log.info('Creating VECTOR INDEX document_chunk_vector_ivf_idx')
cursor.execute(
"""
BEGIN
@@ -318,11 +312,11 @@ class Oracle23aiClient(VectorDBBase):
)
connection.commit()
log.info("Database initialization completed successfully.")
log.info('Database initialization completed successfully.')
except Exception as e:
connection.rollback()
log.exception(f"Error during database initialization: {e}")
log.exception(f'Error during database initialization: {e}')
raise
def check_vector_length(self) -> None:
@@ -344,7 +338,7 @@ class Oracle23aiClient(VectorDBBase):
Returns:
bytes: The vector in Oracle BLOB format
"""
return array.array("f", vector)
return array.array('f', vector)
def adjust_vector_length(self, vector: List[float]) -> List[float]:
"""
@@ -373,7 +367,7 @@ class Oracle23aiClient(VectorDBBase):
"""
if isinstance(obj, Decimal):
return float(obj)
raise TypeError(f"{obj} is not JSON serializable")
raise TypeError(f'{obj} is not JSON serializable')
def _metadata_to_json(self, metadata: Dict) -> str:
"""
@@ -385,7 +379,7 @@ class Oracle23aiClient(VectorDBBase):
Returns:
str: JSON representation of metadata
"""
return json.dumps(metadata, default=self._decimal_handler) if metadata else "{}"
return json.dumps(metadata, default=self._decimal_handler) if metadata else '{}'
def _json_to_metadata(self, json_str: str) -> Dict:
"""
@@ -424,8 +418,8 @@ class Oracle23aiClient(VectorDBBase):
try:
with connection.cursor() as cursor:
for item in items:
vector_blob = self._vector_to_blob(item["vector"])
metadata_json = self._metadata_to_json(item["metadata"])
vector_blob = self._vector_to_blob(item['vector'])
metadata_json = self._metadata_to_json(item['metadata'])
cursor.execute(
"""
@@ -434,22 +428,20 @@ class Oracle23aiClient(VectorDBBase):
VALUES (:id, :collection_name, :text, :metadata, :vector)
""",
{
"id": item["id"],
"collection_name": collection_name,
"text": item["text"],
"metadata": metadata_json,
"vector": vector_blob,
'id': item['id'],
'collection_name': collection_name,
'text': item['text'],
'metadata': metadata_json,
'vector': vector_blob,
},
)
connection.commit()
log.info(
f"Successfully inserted {len(items)} items into collection '{collection_name}'."
)
log.info(f"Successfully inserted {len(items)} items into collection '{collection_name}'.")
except Exception as e:
connection.rollback()
log.exception(f"Error during insert: {e}")
log.exception(f'Error during insert: {e}')
raise
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
@@ -480,8 +472,8 @@ class Oracle23aiClient(VectorDBBase):
try:
with connection.cursor() as cursor:
for item in items:
vector_blob = self._vector_to_blob(item["vector"])
metadata_json = self._metadata_to_json(item["metadata"])
vector_blob = self._vector_to_blob(item['vector'])
metadata_json = self._metadata_to_json(item['metadata'])
cursor.execute(
"""
@@ -499,27 +491,25 @@ class Oracle23aiClient(VectorDBBase):
VALUES (:ins_id, :ins_collection_name, :ins_text, :ins_metadata, :ins_vector)
""",
{
"merge_id": item["id"],
"upd_collection_name": collection_name,
"upd_text": item["text"],
"upd_metadata": metadata_json,
"upd_vector": vector_blob,
"ins_id": item["id"],
"ins_collection_name": collection_name,
"ins_text": item["text"],
"ins_metadata": metadata_json,
"ins_vector": vector_blob,
'merge_id': item['id'],
'upd_collection_name': collection_name,
'upd_text': item['text'],
'upd_metadata': metadata_json,
'upd_vector': vector_blob,
'ins_id': item['id'],
'ins_collection_name': collection_name,
'ins_text': item['text'],
'ins_metadata': metadata_json,
'ins_vector': vector_blob,
},
)
connection.commit()
log.info(
f"Successfully upserted {len(items)} items into collection '{collection_name}'."
)
log.info(f"Successfully upserted {len(items)} items into collection '{collection_name}'.")
except Exception as e:
connection.rollback()
log.exception(f"Error during upsert: {e}")
log.exception(f'Error during upsert: {e}')
raise
def search(
@@ -551,13 +541,11 @@ class Oracle23aiClient(VectorDBBase):
... for i, (id, dist) in enumerate(zip(results.ids[0], results.distances[0])):
... log.info(f"Match {i+1}: id={id}, distance={dist}")
"""
log.info(
f"Searching items from collection '{collection_name}' with limit {limit}."
)
log.info(f"Searching items from collection '{collection_name}' with limit {limit}.")
try:
if not vectors:
log.warning("No vectors provided for search.")
log.warning('No vectors provided for search.')
return None
num_queries = len(vectors)
@@ -583,9 +571,9 @@ class Oracle23aiClient(VectorDBBase):
FETCH APPROX FIRST :limit ROWS ONLY
""",
{
"query_vector": vector_blob,
"collection_name": collection_name,
"limit": limit,
'query_vector': vector_blob,
'collection_name': collection_name,
'limit': limit,
},
)
@@ -593,35 +581,21 @@ class Oracle23aiClient(VectorDBBase):
for row in results:
ids[qid].append(row[0])
documents[qid].append(
row[1].read()
if isinstance(row[1], oracledb.LOB)
else str(row[1])
)
documents[qid].append(row[1].read() if isinstance(row[1], oracledb.LOB) else str(row[1]))
# 🔧 FIXED: Parse JSON metadata properly
metadata_str = (
row[2].read()
if isinstance(row[2], oracledb.LOB)
else row[2]
)
metadata_str = row[2].read() if isinstance(row[2], oracledb.LOB) else row[2]
metadatas[qid].append(self._json_to_metadata(metadata_str))
distances[qid].append(float(row[3]))
log.info(
f"Search completed. Found {sum(len(ids[i]) for i in range(num_queries))} total results."
)
log.info(f'Search completed. Found {sum(len(ids[i]) for i in range(num_queries))} total results.')
return SearchResult(
ids=ids, distances=distances, documents=documents, metadatas=metadatas
)
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
except Exception as e:
log.exception(f"Error during search: {e}")
log.exception(f'Error during search: {e}')
return None
def query(
self, collection_name: str, filter: Dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
"""
Query items based on metadata filters.
@@ -653,15 +627,15 @@ class Oracle23aiClient(VectorDBBase):
WHERE collection_name = :collection_name
"""
params = {"collection_name": collection_name}
params = {'collection_name': collection_name}
for i, (key, value) in enumerate(filter.items()):
param_name = f"value_{i}"
param_name = f'value_{i}'
query += f" AND JSON_VALUE(vmetadata, '$.{key}' RETURNING VARCHAR2(4096)) = :{param_name}"
params[param_name] = str(value)
query += " FETCH FIRST :limit ROWS ONLY"
params["limit"] = limit
query += ' FETCH FIRST :limit ROWS ONLY'
params['limit'] = limit
with self.get_connection() as connection:
with connection.cursor() as cursor:
@@ -669,32 +643,25 @@ class Oracle23aiClient(VectorDBBase):
results = cursor.fetchall()
if not results:
log.info("No results found for query.")
log.info('No results found for query.')
return None
ids = [[row[0] for row in results]]
documents = [
[
row[1].read() if isinstance(row[1], oracledb.LOB) else str(row[1])
for row in results
]
]
documents = [[row[1].read() if isinstance(row[1], oracledb.LOB) else str(row[1]) for row in results]]
# 🔧 FIXED: Parse JSON metadata properly
metadatas = [
[
self._json_to_metadata(
row[2].read() if isinstance(row[2], oracledb.LOB) else row[2]
)
self._json_to_metadata(row[2].read() if isinstance(row[2], oracledb.LOB) else row[2])
for row in results
]
]
log.info(f"Query completed. Found {len(results)} results.")
log.info(f'Query completed. Found {len(results)} results.')
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
except Exception as e:
log.exception(f"Error during query: {e}")
log.exception(f'Error during query: {e}')
return None
def get(self, collection_name: str) -> Optional[GetResult]:
@@ -729,28 +696,21 @@ class Oracle23aiClient(VectorDBBase):
WHERE collection_name = :collection_name
FETCH FIRST :limit ROWS ONLY
""",
{"collection_name": collection_name, "limit": limit},
{'collection_name': collection_name, 'limit': limit},
)
results = cursor.fetchall()
if not results:
log.info("No results found.")
log.info('No results found.')
return None
ids = [[row[0] for row in results]]
documents = [
[
row[1].read() if isinstance(row[1], oracledb.LOB) else str(row[1])
for row in results
]
]
documents = [[row[1].read() if isinstance(row[1], oracledb.LOB) else str(row[1]) for row in results]]
# 🔧 FIXED: Parse JSON metadata properly
metadatas = [
[
self._json_to_metadata(
row[2].read() if isinstance(row[2], oracledb.LOB) else row[2]
)
self._json_to_metadata(row[2].read() if isinstance(row[2], oracledb.LOB) else row[2])
for row in results
]
]
@@ -758,7 +718,7 @@ class Oracle23aiClient(VectorDBBase):
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
except Exception as e:
log.exception(f"Error during get: {e}")
log.exception(f'Error during get: {e}')
return None
def delete(
@@ -790,21 +750,19 @@ class Oracle23aiClient(VectorDBBase):
log.info(f"Deleting items from collection '{collection_name}'.")
try:
query = (
"DELETE FROM document_chunk WHERE collection_name = :collection_name"
)
params = {"collection_name": collection_name}
query = 'DELETE FROM document_chunk WHERE collection_name = :collection_name'
params = {'collection_name': collection_name}
if ids:
# 🔧 FIXED: Use proper parameterized query to prevent SQL injection
placeholders = ",".join([f":id_{i}" for i in range(len(ids))])
query += f" AND id IN ({placeholders})"
placeholders = ','.join([f':id_{i}' for i in range(len(ids))])
query += f' AND id IN ({placeholders})'
for i, id_val in enumerate(ids):
params[f"id_{i}"] = id_val
params[f'id_{i}'] = id_val
if filter:
for i, (key, value) in enumerate(filter.items()):
param_name = f"value_{i}"
param_name = f'value_{i}'
query += f" AND JSON_VALUE(vmetadata, '$.{key}' RETURNING VARCHAR2(4096)) = :{param_name}"
params[param_name] = str(value)
@@ -817,7 +775,7 @@ class Oracle23aiClient(VectorDBBase):
log.info(f"Deleted {deleted} items from collection '{collection_name}'.")
except Exception as e:
log.exception(f"Error during delete: {e}")
log.exception(f'Error during delete: {e}')
raise
def reset(self) -> None:
@@ -833,21 +791,19 @@ class Oracle23aiClient(VectorDBBase):
>>> client = Oracle23aiClient()
>>> client.reset() # Warning: Removes all data!
"""
log.info("Resetting database - deleting all items.")
log.info('Resetting database - deleting all items.')
try:
with self.get_connection() as connection:
with connection.cursor() as cursor:
cursor.execute("DELETE FROM document_chunk")
cursor.execute('DELETE FROM document_chunk')
deleted = cursor.rowcount
connection.commit()
log.info(
f"Reset complete. Deleted {deleted} items from 'document_chunk' table."
)
log.info(f"Reset complete. Deleted {deleted} items from 'document_chunk' table.")
except Exception as e:
log.exception(f"Error during reset: {e}")
log.exception(f'Error during reset: {e}')
raise
def close(self) -> None:
@@ -862,11 +818,11 @@ class Oracle23aiClient(VectorDBBase):
>>> client.close()
"""
try:
if hasattr(self, "pool") and self.pool:
if hasattr(self, 'pool') and self.pool:
self.pool.close()
log.info("Oracle Vector Search connection pool closed.")
log.info('Oracle Vector Search connection pool closed.')
except Exception as e:
log.exception(f"Error closing connection pool: {e}")
log.exception(f'Error closing connection pool: {e}')
def has_collection(self, collection_name: str) -> bool:
"""
@@ -895,7 +851,7 @@ class Oracle23aiClient(VectorDBBase):
WHERE collection_name = :collection_name
FETCH FIRST 1 ROWS ONLY
""",
{"collection_name": collection_name},
{'collection_name': collection_name},
)
count = cursor.fetchone()[0]
@@ -903,7 +859,7 @@ class Oracle23aiClient(VectorDBBase):
return count > 0
except Exception as e:
log.exception(f"Error checking collection existence: {e}")
log.exception(f'Error checking collection existence: {e}')
return False
def delete_collection(self, collection_name: str) -> None:
@@ -929,15 +885,13 @@ class Oracle23aiClient(VectorDBBase):
DELETE FROM document_chunk
WHERE collection_name = :collection_name
""",
{"collection_name": collection_name},
{'collection_name': collection_name},
)
deleted = cursor.rowcount
connection.commit()
log.info(
f"Collection '{collection_name}' deleted. Removed {deleted} items."
)
log.info(f"Collection '{collection_name}' deleted. Removed {deleted} items.")
except Exception as e:
log.exception(f"Error deleting collection '{collection_name}': {e}")
@@ -55,7 +55,7 @@ VECTOR_LENGTH = PGVECTOR_INITIALIZE_MAX_VECTOR_LENGTH
USE_HALFVEC = PGVECTOR_USE_HALFVEC
VECTOR_TYPE_FACTORY = HALFVEC if USE_HALFVEC else Vector
VECTOR_OPCLASS = "halfvec_cosine_ops" if USE_HALFVEC else "vector_cosine_ops"
VECTOR_OPCLASS = 'halfvec_cosine_ops' if USE_HALFVEC else 'vector_cosine_ops'
Base = declarative_base()
log = logging.getLogger(__name__)
@@ -65,12 +65,12 @@ def pgcrypto_encrypt(val, key):
return func.pgp_sym_encrypt(val, literal(key))
def pgcrypto_decrypt(col, key, outtype="text"):
def pgcrypto_decrypt(col, key, outtype='text'):
return func.cast(func.pgp_sym_decrypt(col, literal(key)), outtype)
class DocumentChunk(Base):
__tablename__ = "document_chunk"
__tablename__ = 'document_chunk'
id = Column(Text, primary_key=True)
vector = Column(VECTOR_TYPE_FACTORY(dim=VECTOR_LENGTH), nullable=True)
@@ -86,7 +86,6 @@ class DocumentChunk(Base):
class PgvectorClient(VectorDBBase):
def __init__(self) -> None:
# if no pgvector uri, use the existing database connection
if not PGVECTOR_DB_URL:
from open_webui.internal.db import ScopedSession
@@ -105,46 +104,44 @@ class PgvectorClient(VectorDBBase):
poolclass=QueuePool,
)
else:
engine = create_engine(
PGVECTOR_DB_URL, pool_pre_ping=True, poolclass=NullPool
)
engine = create_engine(PGVECTOR_DB_URL, pool_pre_ping=True, poolclass=NullPool)
else:
engine = create_engine(PGVECTOR_DB_URL, pool_pre_ping=True)
SessionLocal = sessionmaker(
autocommit=False, autoflush=False, bind=engine, expire_on_commit=False
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
self.session = scoped_session(SessionLocal)
try:
# Ensure the pgvector extension is available
# Use a conditional check to avoid permission issues on Azure PostgreSQL
if PGVECTOR_CREATE_EXTENSION:
self.session.execute(text("""
self.session.execute(
text("""
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_extension WHERE extname = 'vector') THEN
CREATE EXTENSION IF NOT EXISTS vector;
END IF;
END $$;
"""))
""")
)
if PGVECTOR_PGCRYPTO:
# Ensure the pgcrypto extension is available for encryption
# Use a conditional check to avoid permission issues on Azure PostgreSQL
self.session.execute(text("""
self.session.execute(
text("""
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_extension WHERE extname = 'pgcrypto') THEN
CREATE EXTENSION IF NOT EXISTS pgcrypto;
END IF;
END $$;
"""))
""")
)
if not PGVECTOR_PGCRYPTO_KEY:
raise ValueError(
"PGVECTOR_PGCRYPTO_KEY must be set when PGVECTOR_PGCRYPTO is enabled."
)
raise ValueError('PGVECTOR_PGCRYPTO_KEY must be set when PGVECTOR_PGCRYPTO is enabled.')
# Check vector length consistency
self.check_vector_length()
@@ -160,15 +157,14 @@ class PgvectorClient(VectorDBBase):
self.session.execute(
text(
"CREATE INDEX IF NOT EXISTS idx_document_chunk_collection_name "
"ON document_chunk (collection_name);"
'CREATE INDEX IF NOT EXISTS idx_document_chunk_collection_name ON document_chunk (collection_name);'
)
)
self.session.commit()
log.info("Initialization complete.")
log.info('Initialization complete.')
except Exception as e:
self.session.rollback()
log.exception(f"Error during initialization: {e}")
log.exception(f'Error during initialization: {e}')
raise
@staticmethod
@@ -176,7 +172,7 @@ class PgvectorClient(VectorDBBase):
if not index_def:
return None
try:
after_using = index_def.lower().split("using ", 1)[1]
after_using = index_def.lower().split('using ', 1)[1]
return after_using.split()[0]
except (IndexError, AttributeError):
return None
@@ -189,23 +185,23 @@ class PgvectorClient(VectorDBBase):
index_method,
)
elif USE_HALFVEC:
index_method = "hnsw"
index_method = 'hnsw'
log.info(
"VECTOR_LENGTH=%s exceeds 2000; using halfvec column type with hnsw index.",
'VECTOR_LENGTH=%s exceeds 2000; using halfvec column type with hnsw index.',
VECTOR_LENGTH,
)
else:
index_method = "ivfflat"
index_method = 'ivfflat'
if index_method == "hnsw":
index_options = f"WITH (m = {PGVECTOR_HNSW_M}, ef_construction = {PGVECTOR_HNSW_EF_CONSTRUCTION})"
if index_method == 'hnsw':
index_options = f'WITH (m = {PGVECTOR_HNSW_M}, ef_construction = {PGVECTOR_HNSW_EF_CONSTRUCTION})'
else:
index_options = f"WITH (lists = {PGVECTOR_IVFFLAT_LISTS})"
index_options = f'WITH (lists = {PGVECTOR_IVFFLAT_LISTS})'
return index_method, index_options
def _ensure_vector_index(self, index_method: str, index_options: str) -> None:
index_name = "idx_document_chunk_vector"
index_name = 'idx_document_chunk_vector'
existing_index_def = self.session.execute(
text("""
SELECT indexdef
@@ -214,7 +210,7 @@ class PgvectorClient(VectorDBBase):
AND tablename = 'document_chunk'
AND indexname = :index_name
"""),
{"index_name": index_name},
{'index_name': index_name},
).scalar()
existing_method = self._extract_index_method(existing_index_def)
@@ -222,23 +218,23 @@ class PgvectorClient(VectorDBBase):
raise RuntimeError(
f"Existing pgvector index '{index_name}' uses method '{existing_method}' but configuration now "
f"requires '{index_method}'. Automatic rebuild is disabled to prevent long-running maintenance. "
"Drop the index manually (optionally after tuning maintenance_work_mem/max_parallel_maintenance_workers) "
"and recreate it with the new method before restarting Open WebUI."
'Drop the index manually (optionally after tuning maintenance_work_mem/max_parallel_maintenance_workers) '
'and recreate it with the new method before restarting Open WebUI.'
)
if not existing_index_def:
index_sql = (
f"CREATE INDEX IF NOT EXISTS {index_name} "
f"ON document_chunk USING {index_method} (vector {VECTOR_OPCLASS})"
f'CREATE INDEX IF NOT EXISTS {index_name} '
f'ON document_chunk USING {index_method} (vector {VECTOR_OPCLASS})'
)
if index_options:
index_sql = f"{index_sql} {index_options}"
index_sql = f'{index_sql} {index_options}'
self.session.execute(text(index_sql))
log.info(
"Ensured vector index '%s' using %s%s.",
index_name,
index_method,
f" {index_options}" if index_options else "",
f' {index_options}' if index_options else '',
)
def check_vector_length(self) -> None:
@@ -249,16 +245,14 @@ class PgvectorClient(VectorDBBase):
metadata = MetaData()
try:
# Attempt to reflect the 'document_chunk' table
document_chunk_table = Table(
"document_chunk", metadata, autoload_with=self.session.bind
)
document_chunk_table = Table('document_chunk', metadata, autoload_with=self.session.bind)
except NoSuchTableError:
# Table does not exist; no action needed
return
# Proceed to check the vector column
if "vector" in document_chunk_table.columns:
vector_column = document_chunk_table.columns["vector"]
if 'vector' in document_chunk_table.columns:
vector_column = document_chunk_table.columns['vector']
vector_type = vector_column.type
expected_type = HALFVEC if USE_HALFVEC else Vector
@@ -268,16 +262,14 @@ class PgvectorClient(VectorDBBase):
f"('{expected_type.__name__}') for VECTOR_LENGTH {VECTOR_LENGTH}."
)
db_vector_length = getattr(vector_type, "dim", None)
db_vector_length = getattr(vector_type, 'dim', None)
if db_vector_length is not None and db_vector_length != VECTOR_LENGTH:
raise Exception(
f"VECTOR_LENGTH {VECTOR_LENGTH} does not match existing vector column dimension {db_vector_length}. "
"Cannot change vector size after initialization without migrating the data."
f'VECTOR_LENGTH {VECTOR_LENGTH} does not match existing vector column dimension {db_vector_length}. '
'Cannot change vector size after initialization without migrating the data.'
)
else:
raise Exception(
"The 'vector' column does not exist in the 'document_chunk' table."
)
raise Exception("The 'vector' column does not exist in the 'document_chunk' table.")
def adjust_vector_length(self, vector: List[float]) -> List[float]:
# Adjust vector to have length VECTOR_LENGTH
@@ -294,10 +286,10 @@ class PgvectorClient(VectorDBBase):
try:
if PGVECTOR_PGCRYPTO:
for item in items:
vector = self.adjust_vector_length(item["vector"])
vector = self.adjust_vector_length(item['vector'])
# Use raw SQL for BYTEA/pgcrypto
# Ensure metadata is converted to its JSON text representation
json_metadata = json.dumps(item["metadata"])
json_metadata = json.dumps(item['metadata'])
self.session.execute(
text("""
INSERT INTO document_chunk
@@ -310,12 +302,12 @@ class PgvectorClient(VectorDBBase):
ON CONFLICT (id) DO NOTHING
"""),
{
"id": item["id"],
"vector": vector,
"collection_name": collection_name,
"text": item["text"],
"metadata_text": json_metadata,
"key": PGVECTOR_PGCRYPTO_KEY,
'id': item['id'],
'vector': vector,
'collection_name': collection_name,
'text': item['text'],
'metadata_text': json_metadata,
'key': PGVECTOR_PGCRYPTO_KEY,
},
)
self.session.commit()
@@ -324,31 +316,29 @@ class PgvectorClient(VectorDBBase):
else:
new_items = []
for item in items:
vector = self.adjust_vector_length(item["vector"])
vector = self.adjust_vector_length(item['vector'])
new_chunk = DocumentChunk(
id=item["id"],
id=item['id'],
vector=vector,
collection_name=collection_name,
text=item["text"],
vmetadata=process_metadata(item["metadata"]),
text=item['text'],
vmetadata=process_metadata(item['metadata']),
)
new_items.append(new_chunk)
self.session.bulk_save_objects(new_items)
self.session.commit()
log.info(
f"Inserted {len(new_items)} items into collection '{collection_name}'."
)
log.info(f"Inserted {len(new_items)} items into collection '{collection_name}'.")
except Exception as e:
self.session.rollback()
log.exception(f"Error during insert: {e}")
log.exception(f'Error during insert: {e}')
raise
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
try:
if PGVECTOR_PGCRYPTO:
for item in items:
vector = self.adjust_vector_length(item["vector"])
json_metadata = json.dumps(item["metadata"])
vector = self.adjust_vector_length(item['vector'])
json_metadata = json.dumps(item['metadata'])
self.session.execute(
text("""
INSERT INTO document_chunk
@@ -365,47 +355,39 @@ class PgvectorClient(VectorDBBase):
vmetadata = EXCLUDED.vmetadata
"""),
{
"id": item["id"],
"vector": vector,
"collection_name": collection_name,
"text": item["text"],
"metadata_text": json_metadata,
"key": PGVECTOR_PGCRYPTO_KEY,
'id': item['id'],
'vector': vector,
'collection_name': collection_name,
'text': item['text'],
'metadata_text': json_metadata,
'key': PGVECTOR_PGCRYPTO_KEY,
},
)
self.session.commit()
log.info(f"Encrypted & upserted {len(items)} into '{collection_name}'")
else:
for item in items:
vector = self.adjust_vector_length(item["vector"])
existing = (
self.session.query(DocumentChunk)
.filter(DocumentChunk.id == item["id"])
.first()
)
vector = self.adjust_vector_length(item['vector'])
existing = self.session.query(DocumentChunk).filter(DocumentChunk.id == item['id']).first()
if existing:
existing.vector = vector
existing.text = item["text"]
existing.vmetadata = process_metadata(item["metadata"])
existing.collection_name = (
collection_name # Update collection_name if necessary
)
existing.text = item['text']
existing.vmetadata = process_metadata(item['metadata'])
existing.collection_name = collection_name # Update collection_name if necessary
else:
new_chunk = DocumentChunk(
id=item["id"],
id=item['id'],
vector=vector,
collection_name=collection_name,
text=item["text"],
vmetadata=process_metadata(item["metadata"]),
text=item['text'],
vmetadata=process_metadata(item['metadata']),
)
self.session.add(new_chunk)
self.session.commit()
log.info(
f"Upserted {len(items)} items into collection '{collection_name}'."
)
log.info(f"Upserted {len(items)} items into collection '{collection_name}'.")
except Exception as e:
self.session.rollback()
log.exception(f"Error during upsert: {e}")
log.exception(f'Error during upsert: {e}')
raise
def search(
@@ -427,38 +409,26 @@ class PgvectorClient(VectorDBBase):
return cast(array(vector), VECTOR_TYPE_FACTORY(VECTOR_LENGTH))
# Create the values for query vectors
qid_col = column("qid", Integer)
q_vector_col = column("q_vector", VECTOR_TYPE_FACTORY(VECTOR_LENGTH))
qid_col = column('qid', Integer)
q_vector_col = column('q_vector', VECTOR_TYPE_FACTORY(VECTOR_LENGTH))
query_vectors = (
values(qid_col, q_vector_col)
.data(
[(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)]
)
.alias("query_vectors")
.data([(idx, vector_expr(vector)) for idx, vector in enumerate(vectors)])
.alias('query_vectors')
)
result_fields = [
DocumentChunk.id,
]
if PGVECTOR_PGCRYPTO:
result_fields.append(pgcrypto_decrypt(DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text).label('text'))
result_fields.append(
pgcrypto_decrypt(
DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text
).label("text")
)
result_fields.append(
pgcrypto_decrypt(
DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB
).label("vmetadata")
pgcrypto_decrypt(DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB).label('vmetadata')
)
else:
result_fields.append(DocumentChunk.text)
result_fields.append(DocumentChunk.vmetadata)
result_fields.append(
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label(
"distance"
)
)
result_fields.append((DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)).label('distance'))
# Build the lateral subquery for each query vector
where_clauses = [DocumentChunk.collection_name == collection_name]
@@ -466,9 +436,9 @@ class PgvectorClient(VectorDBBase):
# Apply metadata filter if provided
if filter:
for key, value in filter.items():
if isinstance(value, dict) and "$in" in value:
if isinstance(value, dict) and '$in' in value:
# Handle $in operator: {"field": {"$in": [values]}}
in_values = value["$in"]
in_values = value['$in']
if PGVECTOR_PGCRYPTO:
where_clauses.append(
pgcrypto_decrypt(
@@ -478,11 +448,7 @@ class PgvectorClient(VectorDBBase):
)[key].astext.in_([str(v) for v in in_values])
)
else:
where_clauses.append(
DocumentChunk.vmetadata[key].astext.in_(
[str(v) for v in in_values]
)
)
where_clauses.append(DocumentChunk.vmetadata[key].astext.in_([str(v) for v in in_values]))
else:
# Handle simple equality: {"field": "value"}
if PGVECTOR_PGCRYPTO:
@@ -495,20 +461,16 @@ class PgvectorClient(VectorDBBase):
== str(value)
)
else:
where_clauses.append(
DocumentChunk.vmetadata[key].astext == str(value)
)
where_clauses.append(DocumentChunk.vmetadata[key].astext == str(value))
subq = (
select(*result_fields)
.where(*where_clauses)
.order_by(
(DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector))
)
.order_by((DocumentChunk.vector.cosine_distance(query_vectors.c.q_vector)))
)
if limit is not None:
subq = subq.limit(limit)
subq = subq.lateral("result")
subq = subq.lateral('result')
# Build the main query by joining query_vectors and the lateral subquery
stmt = (
@@ -550,17 +512,13 @@ class PgvectorClient(VectorDBBase):
metadatas[qid].append(row.vmetadata)
self.session.rollback() # read-only transaction
return SearchResult(
ids=ids, distances=distances, documents=documents, metadatas=metadatas
)
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
except Exception as e:
self.session.rollback()
log.exception(f"Error during search: {e}")
log.exception(f'Error during search: {e}')
return None
def query(
self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
try:
if PGVECTOR_PGCRYPTO:
# Build where clause for vmetadata filter
@@ -568,32 +526,22 @@ class PgvectorClient(VectorDBBase):
for key, value in filter.items():
# decrypt then check key: JSON filter after decryption
where_clauses.append(
pgcrypto_decrypt(
DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB
)[key].astext
pgcrypto_decrypt(DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB)[key].astext
== str(value)
)
stmt = select(
DocumentChunk.id,
pgcrypto_decrypt(
DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text
).label("text"),
pgcrypto_decrypt(
DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB
).label("vmetadata"),
pgcrypto_decrypt(DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text).label('text'),
pgcrypto_decrypt(DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB).label('vmetadata'),
).where(*where_clauses)
if limit is not None:
stmt = stmt.limit(limit)
results = self.session.execute(stmt).all()
else:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
for key, value in filter.items():
query = query.filter(
DocumentChunk.vmetadata[key].astext == str(value)
)
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
if limit is not None:
query = query.limit(limit)
@@ -615,22 +563,16 @@ class PgvectorClient(VectorDBBase):
)
except Exception as e:
self.session.rollback()
log.exception(f"Error during query: {e}")
log.exception(f'Error during query: {e}')
return None
def get(
self, collection_name: str, limit: Optional[int] = None
) -> Optional[GetResult]:
def get(self, collection_name: str, limit: Optional[int] = None) -> Optional[GetResult]:
try:
if PGVECTOR_PGCRYPTO:
stmt = select(
DocumentChunk.id,
pgcrypto_decrypt(
DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text
).label("text"),
pgcrypto_decrypt(
DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB
).label("vmetadata"),
pgcrypto_decrypt(DocumentChunk.text, PGVECTOR_PGCRYPTO_KEY, Text).label('text'),
pgcrypto_decrypt(DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB).label('vmetadata'),
).where(DocumentChunk.collection_name == collection_name)
if limit is not None:
stmt = stmt.limit(limit)
@@ -639,10 +581,7 @@ class PgvectorClient(VectorDBBase):
documents = [[row.text for row in results]]
metadatas = [[row.vmetadata for row in results]]
else:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
if limit is not None:
query = query.limit(limit)
@@ -659,7 +598,7 @@ class PgvectorClient(VectorDBBase):
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
except Exception as e:
self.session.rollback()
log.exception(f"Error during get: {e}")
log.exception(f'Error during get: {e}')
return None
def delete(
@@ -676,43 +615,35 @@ class PgvectorClient(VectorDBBase):
if filter:
for key, value in filter.items():
wheres.append(
pgcrypto_decrypt(
DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB
)[key].astext
pgcrypto_decrypt(DocumentChunk.vmetadata, PGVECTOR_PGCRYPTO_KEY, JSONB)[key].astext
== str(value)
)
stmt = DocumentChunk.__table__.delete().where(*wheres)
result = self.session.execute(stmt)
deleted = result.rowcount
else:
query = self.session.query(DocumentChunk).filter(
DocumentChunk.collection_name == collection_name
)
query = self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name)
if ids:
query = query.filter(DocumentChunk.id.in_(ids))
if filter:
for key, value in filter.items():
query = query.filter(
DocumentChunk.vmetadata[key].astext == str(value)
)
query = query.filter(DocumentChunk.vmetadata[key].astext == str(value))
deleted = query.delete(synchronize_session=False)
self.session.commit()
log.info(f"Deleted {deleted} items from collection '{collection_name}'.")
except Exception as e:
self.session.rollback()
log.exception(f"Error during delete: {e}")
log.exception(f'Error during delete: {e}')
raise
def reset(self) -> None:
try:
deleted = self.session.query(DocumentChunk).delete()
self.session.commit()
log.info(
f"Reset complete. Deleted {deleted} items from 'document_chunk' table."
)
log.info(f"Reset complete. Deleted {deleted} items from 'document_chunk' table.")
except Exception as e:
self.session.rollback()
log.exception(f"Error during reset: {e}")
log.exception(f'Error during reset: {e}')
raise
def close(self) -> None:
@@ -721,16 +652,14 @@ class PgvectorClient(VectorDBBase):
def has_collection(self, collection_name: str) -> bool:
try:
exists = (
self.session.query(DocumentChunk)
.filter(DocumentChunk.collection_name == collection_name)
.first()
self.session.query(DocumentChunk).filter(DocumentChunk.collection_name == collection_name).first()
is not None
)
self.session.rollback() # read-only transaction
return exists
except Exception as e:
self.session.rollback()
log.exception(f"Error checking collection existence: {e}")
log.exception(f'Error checking collection existence: {e}')
return False
def delete_collection(self, collection_name: str) -> None:
@@ -45,7 +45,7 @@ log = logging.getLogger(__name__)
class PineconeClient(VectorDBBase):
def __init__(self):
self.collection_prefix = "open-webui"
self.collection_prefix = 'open-webui'
# Validate required configuration
self._validate_config()
@@ -67,7 +67,7 @@ class PineconeClient(VectorDBBase):
timeout=30, # Reasonable timeout for operations
)
self.using_grpc = True
log.info("Using Pinecone gRPC client for optimal performance")
log.info('Using Pinecone gRPC client for optimal performance')
else:
# Fallback to HTTP client with enhanced connection pooling
self.client = Pinecone(
@@ -76,7 +76,7 @@ class PineconeClient(VectorDBBase):
timeout=30, # Reasonable timeout for operations
)
self.using_grpc = False
log.info("Using Pinecone HTTP client (gRPC not available)")
log.info('Using Pinecone HTTP client (gRPC not available)')
# Persistent executor for batch operations
self._executor = concurrent.futures.ThreadPoolExecutor(max_workers=5)
@@ -88,20 +88,18 @@ class PineconeClient(VectorDBBase):
"""Validate that all required configuration variables are set."""
missing_vars = []
if not PINECONE_API_KEY:
missing_vars.append("PINECONE_API_KEY")
missing_vars.append('PINECONE_API_KEY')
if not PINECONE_ENVIRONMENT:
missing_vars.append("PINECONE_ENVIRONMENT")
missing_vars.append('PINECONE_ENVIRONMENT')
if not PINECONE_INDEX_NAME:
missing_vars.append("PINECONE_INDEX_NAME")
missing_vars.append('PINECONE_INDEX_NAME')
if not PINECONE_DIMENSION:
missing_vars.append("PINECONE_DIMENSION")
missing_vars.append('PINECONE_DIMENSION')
if not PINECONE_CLOUD:
missing_vars.append("PINECONE_CLOUD")
missing_vars.append('PINECONE_CLOUD')
if missing_vars:
raise ValueError(
f"Required configuration missing: {', '.join(missing_vars)}"
)
raise ValueError(f'Required configuration missing: {", ".join(missing_vars)}')
def _initialize_index(self) -> None:
"""Initialize the Pinecone index."""
@@ -126,8 +124,8 @@ class PineconeClient(VectorDBBase):
)
except Exception as e:
log.error(f"Failed to initialize Pinecone index: {e}")
raise RuntimeError(f"Failed to initialize Pinecone index: {e}")
log.error(f'Failed to initialize Pinecone index: {e}')
raise RuntimeError(f'Failed to initialize Pinecone index: {e}')
def _retry_pinecone_operation(self, operation_func, max_retries=3):
"""Retry Pinecone operations with exponential backoff for rate limits and network issues."""
@@ -140,18 +138,18 @@ class PineconeClient(VectorDBBase):
is_retryable = any(
keyword in error_str
for keyword in [
"rate limit",
"quota",
"timeout",
"network",
"connection",
"unavailable",
"internal error",
"429",
"500",
"502",
"503",
"504",
'rate limit',
'quota',
'timeout',
'network',
'connection',
'unavailable',
'internal error',
'429',
'500',
'502',
'503',
'504',
]
)
@@ -162,45 +160,42 @@ class PineconeClient(VectorDBBase):
# Exponential backoff with jitter
delay = (2**attempt) + random.uniform(0, 1)
log.warning(
f"Pinecone operation failed (attempt {attempt + 1}/{max_retries}), "
f"retrying in {delay:.2f}s: {e}"
f'Pinecone operation failed (attempt {attempt + 1}/{max_retries}), retrying in {delay:.2f}s: {e}'
)
time.sleep(delay)
def _create_points(
self, items: List[VectorItem], collection_name_with_prefix: str
) -> List[Dict[str, Any]]:
def _create_points(self, items: List[VectorItem], collection_name_with_prefix: str) -> List[Dict[str, Any]]:
"""Convert VectorItem objects to Pinecone point format."""
points = []
for item in items:
# Start with any existing metadata or an empty dict
metadata = item.get("metadata", {}).copy() if item.get("metadata") else {}
metadata = item.get('metadata', {}).copy() if item.get('metadata') else {}
# Add text to metadata if available
if "text" in item:
metadata["text"] = item["text"]
if 'text' in item:
metadata['text'] = item['text']
# Always add collection_name to metadata for filtering
metadata["collection_name"] = collection_name_with_prefix
metadata['collection_name'] = collection_name_with_prefix
point = {
"id": item["id"],
"values": item["vector"],
"metadata": process_metadata(metadata),
'id': item['id'],
'values': item['vector'],
'metadata': process_metadata(metadata),
}
points.append(point)
return points
def _get_collection_name_with_prefix(self, collection_name: str) -> str:
"""Get the collection name with prefix."""
return f"{self.collection_prefix}_{collection_name}"
return f'{self.collection_prefix}_{collection_name}'
def _normalize_distance(self, score: float) -> float:
"""Normalize distance score based on the metric used."""
if self.metric.lower() == "cosine":
if self.metric.lower() == 'cosine':
# Cosine similarity ranges from -1 to 1, normalize to 0 to 1
return (score + 1.0) / 2.0
elif self.metric.lower() in ["euclidean", "dotproduct"]:
elif self.metric.lower() in ['euclidean', 'dotproduct']:
# These are already suitable for ranking (smaller is better for Euclidean)
return score
else:
@@ -214,68 +209,56 @@ class PineconeClient(VectorDBBase):
metadatas = []
for match in matches:
metadata = getattr(match, "metadata", {}) or {}
ids.append(match.id if hasattr(match, "id") else match["id"])
documents.append(metadata.get("text", ""))
metadata = getattr(match, 'metadata', {}) or {}
ids.append(match.id if hasattr(match, 'id') else match['id'])
documents.append(metadata.get('text', ''))
metadatas.append(metadata)
return GetResult(
**{
"ids": [ids],
"documents": [documents],
"metadatas": [metadatas],
'ids': [ids],
'documents': [documents],
'metadatas': [metadatas],
}
)
def has_collection(self, collection_name: str) -> bool:
"""Check if a collection exists by searching for at least one item."""
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
try:
# Search for at least 1 item with this collection name in metadata
response = self.index.query(
vector=[0.0] * self.dimension, # dummy vector
top_k=1,
filter={"collection_name": collection_name_with_prefix},
filter={'collection_name': collection_name_with_prefix},
include_metadata=False,
)
matches = getattr(response, "matches", []) or []
matches = getattr(response, 'matches', []) or []
return len(matches) > 0
except Exception as e:
log.exception(
f"Error checking collection '{collection_name_with_prefix}': {e}"
)
log.exception(f"Error checking collection '{collection_name_with_prefix}': {e}")
return False
def delete_collection(self, collection_name: str) -> None:
"""Delete a collection by removing all vectors with the collection name in metadata."""
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
try:
self.index.delete(filter={"collection_name": collection_name_with_prefix})
log.info(
f"Collection '{collection_name_with_prefix}' deleted (all vectors removed)."
)
self.index.delete(filter={'collection_name': collection_name_with_prefix})
log.info(f"Collection '{collection_name_with_prefix}' deleted (all vectors removed).")
except Exception as e:
log.warning(
f"Failed to delete collection '{collection_name_with_prefix}': {e}"
)
log.warning(f"Failed to delete collection '{collection_name_with_prefix}': {e}")
raise
def insert(self, collection_name: str, items: List[VectorItem]) -> None:
"""Insert vectors into a collection."""
if not items:
log.warning("No items to insert")
log.warning('No items to insert')
return
start_time = time.time()
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
points = self._create_points(items, collection_name_with_prefix)
# Parallelize batch inserts for performance
@@ -288,26 +271,23 @@ class PineconeClient(VectorDBBase):
try:
future.result()
except Exception as e:
log.error(f"Error inserting batch: {e}")
log.error(f'Error inserting batch: {e}')
raise
elapsed = time.time() - start_time
log.debug(f"Insert of {len(points)} vectors took {elapsed:.2f} seconds")
log.debug(f'Insert of {len(points)} vectors took {elapsed:.2f} seconds')
log.info(
f"Successfully inserted {len(points)} vectors in parallel batches "
f"into '{collection_name_with_prefix}'"
f"Successfully inserted {len(points)} vectors in parallel batches into '{collection_name_with_prefix}'"
)
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
"""Upsert (insert or update) vectors into a collection."""
if not items:
log.warning("No items to upsert")
log.warning('No items to upsert')
return
start_time = time.time()
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
points = self._create_points(items, collection_name_with_prefix)
# Parallelize batch upserts for performance
@@ -320,78 +300,53 @@ class PineconeClient(VectorDBBase):
try:
future.result()
except Exception as e:
log.error(f"Error upserting batch: {e}")
log.error(f'Error upserting batch: {e}')
raise
elapsed = time.time() - start_time
log.debug(f"Upsert of {len(points)} vectors took {elapsed:.2f} seconds")
log.debug(f'Upsert of {len(points)} vectors took {elapsed:.2f} seconds')
log.info(
f"Successfully upserted {len(points)} vectors in parallel batches "
f"into '{collection_name_with_prefix}'"
f"Successfully upserted {len(points)} vectors in parallel batches into '{collection_name_with_prefix}'"
)
async def insert_async(self, collection_name: str, items: List[VectorItem]) -> None:
"""Async version of insert using asyncio and run_in_executor for improved performance."""
if not items:
log.warning("No items to insert")
log.warning('No items to insert')
return
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
points = self._create_points(items, collection_name_with_prefix)
# Create batches
batches = [
points[i : i + BATCH_SIZE] for i in range(0, len(points), BATCH_SIZE)
]
batches = [points[i : i + BATCH_SIZE] for i in range(0, len(points), BATCH_SIZE)]
loop = asyncio.get_event_loop()
tasks = [
loop.run_in_executor(
None, functools.partial(self.index.upsert, vectors=batch)
)
for batch in batches
]
tasks = [loop.run_in_executor(None, functools.partial(self.index.upsert, vectors=batch)) for batch in batches]
results = await asyncio.gather(*tasks, return_exceptions=True)
for result in results:
if isinstance(result, Exception):
log.error(f"Error in async insert batch: {result}")
log.error(f'Error in async insert batch: {result}')
raise result
log.info(
f"Successfully async inserted {len(points)} vectors in batches "
f"into '{collection_name_with_prefix}'"
)
log.info(f"Successfully async inserted {len(points)} vectors in batches into '{collection_name_with_prefix}'")
async def upsert_async(self, collection_name: str, items: List[VectorItem]) -> None:
"""Async version of upsert using asyncio and run_in_executor for improved performance."""
if not items:
log.warning("No items to upsert")
log.warning('No items to upsert')
return
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
points = self._create_points(items, collection_name_with_prefix)
# Create batches
batches = [
points[i : i + BATCH_SIZE] for i in range(0, len(points), BATCH_SIZE)
]
batches = [points[i : i + BATCH_SIZE] for i in range(0, len(points), BATCH_SIZE)]
loop = asyncio.get_event_loop()
tasks = [
loop.run_in_executor(
None, functools.partial(self.index.upsert, vectors=batch)
)
for batch in batches
]
tasks = [loop.run_in_executor(None, functools.partial(self.index.upsert, vectors=batch)) for batch in batches]
results = await asyncio.gather(*tasks, return_exceptions=True)
for result in results:
if isinstance(result, Exception):
log.error(f"Error in async upsert batch: {result}")
log.error(f'Error in async upsert batch: {result}')
raise result
log.info(
f"Successfully async upserted {len(points)} vectors in batches "
f"into '{collection_name_with_prefix}'"
)
log.info(f"Successfully async upserted {len(points)} vectors in batches into '{collection_name_with_prefix}'")
def search(
self,
@@ -402,12 +357,10 @@ class PineconeClient(VectorDBBase):
) -> Optional[SearchResult]:
"""Search for similar vectors in a collection."""
if not vectors or not vectors[0]:
log.warning("No vectors provided for search")
log.warning('No vectors provided for search')
return None
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
if limit is None or limit <= 0:
limit = NO_LIMIT
@@ -421,10 +374,10 @@ class PineconeClient(VectorDBBase):
vector=query_vector,
top_k=limit,
include_metadata=True,
filter={"collection_name": collection_name_with_prefix},
filter={'collection_name': collection_name_with_prefix},
)
matches = getattr(query_response, "matches", []) or []
matches = getattr(query_response, 'matches', []) or []
if not matches:
# Return empty result if no matches
return SearchResult(
@@ -438,12 +391,7 @@ class PineconeClient(VectorDBBase):
get_result = self._result_to_get_result(matches)
# Calculate normalized distances based on metric
distances = [
[
self._normalize_distance(getattr(match, "score", 0.0))
for match in matches
]
]
distances = [[self._normalize_distance(getattr(match, 'score', 0.0)) for match in matches]]
return SearchResult(
ids=get_result.ids,
@@ -455,13 +403,9 @@ class PineconeClient(VectorDBBase):
log.error(f"Error searching in '{collection_name_with_prefix}': {e}")
return None
def query(
self, collection_name: str, filter: Dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
"""Query vectors by metadata filter."""
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
if limit is None or limit <= 0:
limit = NO_LIMIT
@@ -471,7 +415,7 @@ class PineconeClient(VectorDBBase):
zero_vector = [0.0] * self.dimension
# Combine user filter with collection_name
pinecone_filter = {"collection_name": collection_name_with_prefix}
pinecone_filter = {'collection_name': collection_name_with_prefix}
if filter:
pinecone_filter.update(filter)
@@ -483,7 +427,7 @@ class PineconeClient(VectorDBBase):
include_metadata=True,
)
matches = getattr(query_response, "matches", []) or []
matches = getattr(query_response, 'matches', []) or []
return self._result_to_get_result(matches)
except Exception as e:
@@ -492,9 +436,7 @@ class PineconeClient(VectorDBBase):
def get(self, collection_name: str) -> Optional[GetResult]:
"""Get all vectors in a collection."""
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
try:
# Use a zero vector for fetching all entries
@@ -505,10 +447,10 @@ class PineconeClient(VectorDBBase):
vector=zero_vector,
top_k=NO_LIMIT,
include_metadata=True,
filter={"collection_name": collection_name_with_prefix},
filter={'collection_name': collection_name_with_prefix},
)
matches = getattr(query_response, "matches", []) or []
matches = getattr(query_response, 'matches', []) or []
return self._result_to_get_result(matches)
except Exception as e:
@@ -522,9 +464,7 @@ class PineconeClient(VectorDBBase):
filter: Optional[Dict] = None,
) -> None:
"""Delete vectors by IDs or filter."""
collection_name_with_prefix = self._get_collection_name_with_prefix(
collection_name
)
collection_name_with_prefix = self._get_collection_name_with_prefix(collection_name)
try:
if ids:
@@ -534,28 +474,20 @@ class PineconeClient(VectorDBBase):
# Note: When deleting by ID, we can't filter by collection_name
# This is a limitation of Pinecone - be careful with ID uniqueness
self.index.delete(ids=batch_ids)
log.debug(
f"Deleted batch of {len(batch_ids)} vectors by ID "
f"from '{collection_name_with_prefix}'"
)
log.info(
f"Successfully deleted {len(ids)} vectors by ID "
f"from '{collection_name_with_prefix}'"
)
log.debug(f"Deleted batch of {len(batch_ids)} vectors by ID from '{collection_name_with_prefix}'")
log.info(f"Successfully deleted {len(ids)} vectors by ID from '{collection_name_with_prefix}'")
elif filter:
# Combine user filter with collection_name
pinecone_filter = {"collection_name": collection_name_with_prefix}
pinecone_filter = {'collection_name': collection_name_with_prefix}
if filter:
pinecone_filter.update(filter)
# Delete by metadata filter
self.index.delete(filter=pinecone_filter)
log.info(
f"Successfully deleted vectors by filter from '{collection_name_with_prefix}'"
)
log.info(f"Successfully deleted vectors by filter from '{collection_name_with_prefix}'")
else:
log.warning("No ids or filter provided for delete operation")
log.warning('No ids or filter provided for delete operation')
except Exception as e:
log.error(f"Error deleting from collection '{collection_name}': {e}")
@@ -565,9 +497,9 @@ class PineconeClient(VectorDBBase):
"""Reset the database by deleting all collections."""
try:
self.index.delete(delete_all=True)
log.info("All vectors successfully deleted from the index.")
log.info('All vectors successfully deleted from the index.')
except Exception as e:
log.error(f"Failed to reset Pinecone index: {e}")
log.error(f'Failed to reset Pinecone index: {e}')
raise
def close(self):
@@ -576,7 +508,7 @@ class PineconeClient(VectorDBBase):
# The new Pinecone client doesn't need explicit closing
pass
except Exception as e:
log.warning(f"Failed to clean up Pinecone resources: {e}")
log.warning(f'Failed to clean up Pinecone resources: {e}')
self._executor.shutdown(wait=True)
def __enter__(self):
@@ -76,19 +76,19 @@ class QdrantClient(VectorDBBase):
for point in points:
payload = point.payload
ids.append(point.id)
documents.append(payload["text"])
metadatas.append(payload["metadata"])
documents.append(payload['text'])
metadatas.append(payload['metadata'])
return GetResult(
**{
"ids": [ids],
"documents": [documents],
"metadatas": [metadatas],
'ids': [ids],
'documents': [documents],
'metadatas': [metadatas],
}
)
def _create_collection(self, collection_name: str, dimension: int):
collection_name_with_prefix = f"{self.collection_prefix}_{collection_name}"
collection_name_with_prefix = f'{self.collection_prefix}_{collection_name}'
self.client.create_collection(
collection_name=collection_name_with_prefix,
vectors_config=models.VectorParams(
@@ -104,7 +104,7 @@ class QdrantClient(VectorDBBase):
# Create payload indexes for efficient filtering
self.client.create_payload_index(
collection_name=collection_name_with_prefix,
field_name="metadata.hash",
field_name='metadata.hash',
field_schema=models.KeywordIndexParams(
type=models.KeywordIndexType.KEYWORD,
is_tenant=False,
@@ -113,40 +113,34 @@ class QdrantClient(VectorDBBase):
)
self.client.create_payload_index(
collection_name=collection_name_with_prefix,
field_name="metadata.file_id",
field_name='metadata.file_id',
field_schema=models.KeywordIndexParams(
type=models.KeywordIndexType.KEYWORD,
is_tenant=False,
on_disk=self.QDRANT_ON_DISK,
),
)
log.info(f"collection {collection_name_with_prefix} successfully created!")
log.info(f'collection {collection_name_with_prefix} successfully created!')
def _create_collection_if_not_exists(self, collection_name, dimension):
if not self.has_collection(collection_name=collection_name):
self._create_collection(
collection_name=collection_name, dimension=dimension
)
self._create_collection(collection_name=collection_name, dimension=dimension)
def _create_points(self, items: list[VectorItem]):
return [
PointStruct(
id=item["id"],
vector=item["vector"],
payload={"text": item["text"], "metadata": item["metadata"]},
id=item['id'],
vector=item['vector'],
payload={'text': item['text'], 'metadata': item['metadata']},
)
for item in items
]
def has_collection(self, collection_name: str) -> bool:
return self.client.collection_exists(
f"{self.collection_prefix}_{collection_name}"
)
return self.client.collection_exists(f'{self.collection_prefix}_{collection_name}')
def delete_collection(self, collection_name: str):
return self.client.delete_collection(
collection_name=f"{self.collection_prefix}_{collection_name}"
)
return self.client.delete_collection(collection_name=f'{self.collection_prefix}_{collection_name}')
def search(
self,
@@ -160,7 +154,7 @@ class QdrantClient(VectorDBBase):
limit = NO_LIMIT # otherwise qdrant would set limit to 10!
query_response = self.client.query_points(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
query=vectors[0],
limit=limit,
)
@@ -184,13 +178,11 @@ class QdrantClient(VectorDBBase):
field_conditions = []
for key, value in filter.items():
field_conditions.append(
models.FieldCondition(
key=f"metadata.{key}", match=models.MatchValue(value=value)
)
models.FieldCondition(key=f'metadata.{key}', match=models.MatchValue(value=value))
)
points = self.client.scroll(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
scroll_filter=models.Filter(should=field_conditions),
limit=limit,
)
@@ -202,22 +194,22 @@ class QdrantClient(VectorDBBase):
def get(self, collection_name: str) -> Optional[GetResult]:
# Get all the items in the collection.
points = self.client.scroll(
collection_name=f"{self.collection_prefix}_{collection_name}",
collection_name=f'{self.collection_prefix}_{collection_name}',
limit=NO_LIMIT, # otherwise qdrant would set limit to 10!
)
return self._result_to_get_result(points[0])
def insert(self, collection_name: str, items: list[VectorItem]):
# Insert the items into the collection, if the collection does not exist, it will be created.
self._create_collection_if_not_exists(collection_name, len(items[0]["vector"]))
self._create_collection_if_not_exists(collection_name, len(items[0]['vector']))
points = self._create_points(items)
self.client.upload_points(f"{self.collection_prefix}_{collection_name}", points)
self.client.upload_points(f'{self.collection_prefix}_{collection_name}', points)
def upsert(self, collection_name: str, items: list[VectorItem]):
# Update the items in the collection, if the items are not present, insert them. If the collection does not exist, it will be created.
self._create_collection_if_not_exists(collection_name, len(items[0]["vector"]))
self._create_collection_if_not_exists(collection_name, len(items[0]['vector']))
points = self._create_points(items)
return self.client.upsert(f"{self.collection_prefix}_{collection_name}", points)
return self.client.upsert(f'{self.collection_prefix}_{collection_name}', points)
def delete(
self,
@@ -230,26 +222,28 @@ class QdrantClient(VectorDBBase):
if ids:
for id_value in ids:
field_conditions.append(
models.FieldCondition(
key="metadata.id",
match=models.MatchValue(value=id_value),
(
field_conditions.append(
models.FieldCondition(
key='metadata.id',
match=models.MatchValue(value=id_value),
),
),
),
)
elif filter:
for key, value in filter.items():
field_conditions.append(
models.FieldCondition(
key=f"metadata.{key}",
match=models.MatchValue(value=value),
(
field_conditions.append(
models.FieldCondition(
key=f'metadata.{key}',
match=models.MatchValue(value=value),
),
),
),
)
return self.client.delete(
collection_name=f"{self.collection_prefix}_{collection_name}",
points_selector=models.FilterSelector(
filter=models.Filter(must=field_conditions)
),
collection_name=f'{self.collection_prefix}_{collection_name}',
points_selector=models.FilterSelector(filter=models.Filter(must=field_conditions)),
)
def reset(self):
@@ -29,22 +29,18 @@ from qdrant_client.http.models import PointStruct
from qdrant_client.models import models
NO_LIMIT = 999999999
TENANT_ID_FIELD = "tenant_id"
TENANT_ID_FIELD = 'tenant_id'
DEFAULT_DIMENSION = 384
log = logging.getLogger(__name__)
def _tenant_filter(tenant_id: str) -> models.FieldCondition:
return models.FieldCondition(
key=TENANT_ID_FIELD, match=models.MatchValue(value=tenant_id)
)
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)
)
return models.FieldCondition(key=f'metadata.{key}', match=models.MatchValue(value=value))
class QdrantClient(VectorDBBase):
@@ -59,9 +55,7 @@ class QdrantClient(VectorDBBase):
self.QDRANT_HNSW_M = QDRANT_HNSW_M
if not self.QDRANT_URI:
raise ValueError(
"QDRANT_URI is not set. Please configure it in the environment variables."
)
raise ValueError('QDRANT_URI is not set. Please configure it in the environment variables.')
# Unified handling for either scheme
parsed = urlparse(self.QDRANT_URI)
@@ -86,19 +80,19 @@ class QdrantClient(VectorDBBase):
)
# 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.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'
def _result_to_get_result(self, points) -> GetResult:
ids, documents, metadatas = [], [], []
for point in points:
payload = point.payload
ids.append(point.id)
documents.append(payload["text"])
metadatas.append(payload["metadata"])
documents.append(payload['text'])
metadatas.append(payload['metadata'])
return GetResult(ids=[ids], documents=[documents], metadatas=[metadatas])
def _get_collection_and_tenant_id(self, collection_name: str) -> Tuple[str, str]:
@@ -118,29 +112,25 @@ class QdrantClient(VectorDBBase):
# Check for user memory collections
tenant_id = collection_name
if collection_name.startswith("user-memory-"):
if collection_name.startswith('user-memory-'):
return self.MEMORY_COLLECTION, tenant_id
# Check for file collections
elif collection_name.startswith("file-"):
elif collection_name.startswith('file-'):
return self.FILE_COLLECTION, tenant_id
# Check for web search collections
elif collection_name.startswith("web-search-"):
elif collection_name.startswith('web-search-'):
return self.WEB_SEARCH_COLLECTION, tenant_id
# Handle hash-based collections (YouTube and web URLs)
elif len(collection_name) == 63 and all(
c in "0123456789abcdef" for c in collection_name
):
elif len(collection_name) == 63 and all(c in '0123456789abcdef' for c in collection_name):
return self.HASH_BASED_COLLECTION, tenant_id
else:
return self.KNOWLEDGE_COLLECTION, tenant_id
def _create_multi_tenant_collection(
self, mt_collection_name: str, dimension: int = DEFAULT_DIMENSION
):
def _create_multi_tenant_collection(self, mt_collection_name: str, dimension: int = DEFAULT_DIMENSION):
"""
Creates a collection with multi-tenancy configuration and payload indexes for tenant_id and metadata fields.
"""
@@ -158,9 +148,7 @@ class QdrantClient(VectorDBBase):
m=0,
),
)
log.info(
f"Multi-tenant collection {mt_collection_name} created with dimension {dimension}!"
)
log.info(f'Multi-tenant collection {mt_collection_name} created with dimension {dimension}!')
self.client.create_payload_index(
collection_name=mt_collection_name,
@@ -172,7 +160,7 @@ class QdrantClient(VectorDBBase):
),
)
for field in ("metadata.hash", "metadata.file_id"):
for field in ('metadata.hash', 'metadata.file_id'):
self.client.create_payload_index(
collection_name=mt_collection_name,
field_name=field,
@@ -182,28 +170,24 @@ class QdrantClient(VectorDBBase):
),
)
def _create_points(
self, items: List[VectorItem], tenant_id: str
) -> List[PointStruct]:
def _create_points(self, items: List[VectorItem], tenant_id: str) -> List[PointStruct]:
"""
Create point structs from vector items with tenant ID.
"""
return [
PointStruct(
id=item["id"],
vector=item["vector"],
id=item['id'],
vector=item['vector'],
payload={
"text": item["text"],
"metadata": item["metadata"],
'text': item['text'],
'metadata': item['metadata'],
TENANT_ID_FIELD: tenant_id,
},
)
for item in items
]
def _ensure_collection(
self, mt_collection_name: str, dimension: int = DEFAULT_DIMENSION
):
def _ensure_collection(self, mt_collection_name: str, dimension: int = DEFAULT_DIMENSION):
"""
Ensure the collection exists and payload indexes are created for tenant_id and metadata fields.
"""
@@ -246,15 +230,13 @@ class QdrantClient(VectorDBBase):
must_conditions = [_tenant_filter(tenant_id)]
should_conditions = []
if ids:
should_conditions = [_metadata_filter("id", id_value) for id_value in ids]
should_conditions = [_metadata_filter('id', id_value) for id_value in ids]
elif filter:
must_conditions += [_metadata_filter(k, v) for k, v in filter.items()]
return self.client.delete(
collection_name=mt_collection,
points_selector=models.FilterSelector(
filter=models.Filter(must=must_conditions, should=should_conditions)
),
points_selector=models.FilterSelector(filter=models.Filter(must=must_conditions, should=should_conditions)),
)
def search(
@@ -289,9 +271,7 @@ class QdrantClient(VectorDBBase):
distances=[[(point.score + 1.0) / 2.0 for point in query_response.points]],
)
def query(
self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None
):
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None):
"""
Query points with filters and tenant isolation.
"""
@@ -338,7 +318,7 @@ class QdrantClient(VectorDBBase):
if not self.client or not items:
return None
mt_collection, tenant_id = self._get_collection_and_tenant_id(collection_name)
dimension = len(items[0]["vector"])
dimension = len(items[0]['vector'])
self._ensure_collection(mt_collection, dimension)
points = self._create_points(items, tenant_id)
self.client.upload_points(mt_collection, points)
@@ -372,7 +352,5 @@ class QdrantClient(VectorDBBase):
return None
self.client.delete(
collection_name=mt_collection,
points_selector=models.FilterSelector(
filter=models.Filter(must=[_tenant_filter(tenant_id)])
),
points_selector=models.FilterSelector(filter=models.Filter(must=[_tenant_filter(tenant_id)])),
)
@@ -28,18 +28,16 @@ class S3VectorClient(VectorDBBase):
# Simple validation - log warnings instead of raising exceptions
if not self.bucket_name:
log.warning("S3_VECTOR_BUCKET_NAME not set - S3Vector will not work")
log.warning('S3_VECTOR_BUCKET_NAME not set - S3Vector will not work')
if not self.region:
log.warning("S3_VECTOR_REGION not set - S3Vector will not work")
log.warning('S3_VECTOR_REGION not set - S3Vector will not work')
if self.bucket_name and self.region:
try:
self.client = boto3.client("s3vectors", region_name=self.region)
log.info(
f"S3Vector client initialized for bucket '{self.bucket_name}' in region '{self.region}'"
)
self.client = boto3.client('s3vectors', region_name=self.region)
log.info(f"S3Vector client initialized for bucket '{self.bucket_name}' in region '{self.region}'")
except Exception as e:
log.error(f"Failed to initialize S3Vector client: {e}")
log.error(f'Failed to initialize S3Vector client: {e}')
self.client = None
else:
self.client = None
@@ -48,8 +46,8 @@ class S3VectorClient(VectorDBBase):
self,
index_name: str,
dimension: int,
data_type: str = "float32",
distance_metric: str = "cosine",
data_type: str = 'float32',
distance_metric: str = 'cosine',
) -> None:
"""
Create a new index in the S3 vector bucket for the given collection if it does not exist.
@@ -66,21 +64,17 @@ class S3VectorClient(VectorDBBase):
dimension=dimension,
distanceMetric=distance_metric,
metadataConfiguration={
"nonFilterableMetadataKeys": [
"text",
'nonFilterableMetadataKeys': [
'text',
]
},
)
log.info(
f"Created S3 index: {index_name} (dim={dimension}, type={data_type}, metric={distance_metric})"
)
log.info(f'Created S3 index: {index_name} (dim={dimension}, type={data_type}, metric={distance_metric})')
except Exception as e:
log.error(f"Error creating S3 index '{index_name}': {e}")
raise
def _filter_metadata(
self, metadata: Dict[str, Any], item_id: str
) -> Dict[str, Any]:
def _filter_metadata(self, metadata: Dict[str, Any], item_id: str) -> Dict[str, Any]:
"""
Filter vector metadata keys to comply with S3 Vector API limit of 10 keys maximum.
"""
@@ -89,16 +83,16 @@ class S3VectorClient(VectorDBBase):
# Keep only the first 10 keys, prioritizing important ones based on actual Open WebUI metadata
important_keys = [
"text", # The actual document content
"file_id", # File ID
"source", # Document source file
"title", # Document title
"page", # Page number
"total_pages", # Total pages in document
"embedding_config", # Embedding configuration
"created_by", # User who created it
"name", # Document name
"hash", # Content hash
'text', # The actual document content
'file_id', # File ID
'source', # Document source file
'title', # Document title
'page', # Page number
'total_pages', # Total pages in document
'embedding_config', # Embedding configuration
'created_by', # User who created it
'name', # Document name
'hash', # Content hash
]
filtered_metadata = {}
@@ -117,9 +111,7 @@ class S3VectorClient(VectorDBBase):
if len(filtered_metadata) >= 10:
break
log.warning(
f"Metadata for key '{item_id}' had {len(metadata)} keys, limited to 10 keys"
)
log.warning(f"Metadata for key '{item_id}' had {len(metadata)} keys, limited to 10 keys")
return filtered_metadata
def has_collection(self, collection_name: str) -> bool:
@@ -128,9 +120,7 @@ class S3VectorClient(VectorDBBase):
This avoids pagination issues with list_indexes() and is significantly faster.
"""
try:
self.client.get_index(
vectorBucketName=self.bucket_name, indexName=collection_name
)
self.client.get_index(vectorBucketName=self.bucket_name, indexName=collection_name)
return True
except Exception as e:
log.error(f"Error checking if index '{collection_name}' exists: {e}")
@@ -142,16 +132,12 @@ class S3VectorClient(VectorDBBase):
"""
if not self.has_collection(collection_name):
log.warning(
f"Collection '{collection_name}' does not exist, nothing to delete"
)
log.warning(f"Collection '{collection_name}' does not exist, nothing to delete")
return
try:
log.info(f"Deleting collection '{collection_name}'")
self.client.delete_index(
vectorBucketName=self.bucket_name, indexName=collection_name
)
self.client.delete_index(vectorBucketName=self.bucket_name, indexName=collection_name)
log.info(f"Successfully deleted collection '{collection_name}'")
except Exception as e:
log.error(f"Error deleting collection '{collection_name}': {e}")
@@ -162,10 +148,10 @@ class S3VectorClient(VectorDBBase):
Insert vector items into the S3 Vector index. Create index if it does not exist.
"""
if not items:
log.warning("No items to insert")
log.warning('No items to insert')
return
dimension = len(items[0]["vector"])
dimension = len(items[0]['vector'])
try:
if not self.has_collection(collection_name):
@@ -173,36 +159,36 @@ class S3VectorClient(VectorDBBase):
self._create_index(
index_name=collection_name,
dimension=dimension,
data_type="float32",
distance_metric="cosine",
data_type='float32',
distance_metric='cosine',
)
# Prepare vectors for insertion
vectors = []
for item in items:
# Ensure vector data is in the correct format for S3 Vector API
vector_data = item["vector"]
vector_data = item['vector']
if isinstance(vector_data, list):
# Convert list to float32 values as required by S3 Vector API
vector_data = [float(x) for x in vector_data]
# Prepare metadata, ensuring the text field is preserved
metadata = item.get("metadata", {}).copy()
metadata = item.get('metadata', {}).copy()
# Add the text field to metadata so it's available for retrieval
metadata["text"] = item["text"]
metadata['text'] = item['text']
# Convert metadata to string format for consistency
metadata = process_metadata(metadata)
# Filter metadata to comply with S3 Vector API limit of 10 keys
metadata = self._filter_metadata(metadata, item["id"])
metadata = self._filter_metadata(metadata, item['id'])
vectors.append(
{
"key": item["id"],
"data": {"float32": vector_data},
"metadata": metadata,
'key': item['id'],
'data': {'float32': vector_data},
'metadata': metadata,
}
)
@@ -215,15 +201,11 @@ class S3VectorClient(VectorDBBase):
indexName=collection_name,
vectors=batch,
)
log.info(
f"Inserted batch {i//batch_size + 1}: {len(batch)} vectors into index '{collection_name}'."
)
log.info(f"Inserted batch {i // batch_size + 1}: {len(batch)} vectors into index '{collection_name}'.")
log.info(
f"Completed insertion of {len(vectors)} vectors into index '{collection_name}'."
)
log.info(f"Completed insertion of {len(vectors)} vectors into index '{collection_name}'.")
except Exception as e:
log.error(f"Error inserting vectors: {e}")
log.error(f'Error inserting vectors: {e}')
raise
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
@@ -231,49 +213,47 @@ class S3VectorClient(VectorDBBase):
Insert or update vector items in the S3 Vector index. Create index if it does not exist.
"""
if not items:
log.warning("No items to upsert")
log.warning('No items to upsert')
return
dimension = len(items[0]["vector"])
log.info(f"Upsert dimension: {dimension}")
dimension = len(items[0]['vector'])
log.info(f'Upsert dimension: {dimension}')
try:
if not self.has_collection(collection_name):
log.info(
f"Index '{collection_name}' does not exist. Creating index for upsert."
)
log.info(f"Index '{collection_name}' does not exist. Creating index for upsert.")
self._create_index(
index_name=collection_name,
dimension=dimension,
data_type="float32",
distance_metric="cosine",
data_type='float32',
distance_metric='cosine',
)
# Prepare vectors for upsert
vectors = []
for item in items:
# Ensure vector data is in the correct format for S3 Vector API
vector_data = item["vector"]
vector_data = item['vector']
if isinstance(vector_data, list):
# Convert list to float32 values as required by S3 Vector API
vector_data = [float(x) for x in vector_data]
# Prepare metadata, ensuring the text field is preserved
metadata = item.get("metadata", {}).copy()
metadata = item.get('metadata', {}).copy()
# Add the text field to metadata so it's available for retrieval
metadata["text"] = item["text"]
metadata['text'] = item['text']
# Convert metadata to string format for consistency
metadata = process_metadata(metadata)
# Filter metadata to comply with S3 Vector API limit of 10 keys
metadata = self._filter_metadata(metadata, item["id"])
metadata = self._filter_metadata(metadata, item['id'])
vectors.append(
{
"key": item["id"],
"data": {"float32": vector_data},
"metadata": metadata,
'key': item['id'],
'data': {'float32': vector_data},
'metadata': metadata,
}
)
@@ -283,12 +263,10 @@ class S3VectorClient(VectorDBBase):
batch = vectors[i : i + batch_size]
if i == 0: # Log sample info for first batch only
log.info(
f"Upserting batch 1: {len(batch)} vectors. First vector sample: key={batch[0]['key']}, data_type={type(batch[0]['data']['float32'])}, data_len={len(batch[0]['data']['float32'])}"
f'Upserting batch 1: {len(batch)} vectors. First vector sample: key={batch[0]["key"]}, data_type={type(batch[0]["data"]["float32"])}, data_len={len(batch[0]["data"]["float32"])}'
)
else:
log.info(
f"Upserting batch {i//batch_size + 1}: {len(batch)} vectors."
)
log.info(f'Upserting batch {i // batch_size + 1}: {len(batch)} vectors.')
self.client.put_vectors(
vectorBucketName=self.bucket_name,
@@ -296,11 +274,9 @@ class S3VectorClient(VectorDBBase):
vectors=batch,
)
log.info(
f"Completed upsert of {len(vectors)} vectors into index '{collection_name}'."
)
log.info(f"Completed upsert of {len(vectors)} vectors into index '{collection_name}'.")
except Exception as e:
log.error(f"Error upserting vectors: {e}")
log.error(f'Error upserting vectors: {e}')
raise
def search(
@@ -319,13 +295,11 @@ class S3VectorClient(VectorDBBase):
return None
if not vectors:
log.warning("No query vectors provided")
log.warning('No query vectors provided')
return None
try:
log.info(
f"Searching collection '{collection_name}' with {len(vectors)} query vectors, limit={limit}"
)
log.info(f"Searching collection '{collection_name}' with {len(vectors)} query vectors, limit={limit}")
# Initialize result lists
all_ids = []
@@ -335,10 +309,10 @@ class S3VectorClient(VectorDBBase):
# Process each query vector
for i, query_vector in enumerate(vectors):
log.debug(f"Processing query vector {i+1}/{len(vectors)}")
log.debug(f'Processing query vector {i + 1}/{len(vectors)}')
# Prepare the query vector in S3 Vector format
query_vector_dict = {"float32": [float(x) for x in query_vector]}
query_vector_dict = {'float32': [float(x) for x in query_vector]}
# Call S3 Vector query API
response = self.client.query_vectors(
@@ -356,24 +330,22 @@ class S3VectorClient(VectorDBBase):
query_metadatas = []
query_distances = []
result_vectors = response.get("vectors", [])
result_vectors = response.get('vectors', [])
for vector in result_vectors:
vector_id = vector.get("key")
vector_metadata = vector.get("metadata", {})
vector_distance = vector.get("distance", 0.0)
vector_id = vector.get('key')
vector_metadata = vector.get('metadata', {})
vector_distance = vector.get('distance', 0.0)
# Extract document text from metadata
document_text = ""
document_text = ''
if isinstance(vector_metadata, dict):
# Get the text field first (highest priority)
document_text = vector_metadata.get("text")
document_text = vector_metadata.get('text')
if not document_text:
# Fallback to other possible text fields
document_text = (
vector_metadata.get("content")
or vector_metadata.get("document")
or vector_id
vector_metadata.get('content') or vector_metadata.get('document') or vector_id
)
else:
document_text = vector_id
@@ -389,7 +361,7 @@ class S3VectorClient(VectorDBBase):
all_metadatas.append(query_metadatas)
all_distances.append(query_distances)
log.info(f"Search completed. Found results for {len(all_ids)} queries")
log.info(f'Search completed. Found results for {len(all_ids)} queries')
# Return SearchResult format
return SearchResult(
@@ -402,24 +374,20 @@ class S3VectorClient(VectorDBBase):
except Exception as e:
log.error(f"Error searching collection '{collection_name}': {str(e)}")
# Handle specific AWS exceptions
if hasattr(e, "response") and "Error" in e.response:
error_code = e.response["Error"]["Code"]
if error_code == "NotFoundException":
if hasattr(e, 'response') and 'Error' in e.response:
error_code = e.response['Error']['Code']
if error_code == 'NotFoundException':
log.warning(f"Collection '{collection_name}' not found")
return None
elif error_code == "ValidationException":
log.error(f"Invalid query vector dimensions or parameters")
elif error_code == 'ValidationException':
log.error(f'Invalid query vector dimensions or parameters')
return None
elif error_code == "AccessDeniedException":
log.error(
f"Access denied for collection '{collection_name}'. Check permissions."
)
elif error_code == 'AccessDeniedException':
log.error(f"Access denied for collection '{collection_name}'. Check permissions.")
return None
raise
def query(
self, collection_name: str, filter: Dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
"""
Query vectors from a collection using metadata filter.
"""
@@ -429,7 +397,7 @@ class S3VectorClient(VectorDBBase):
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
if not filter:
log.warning("No filter provided, returning all vectors")
log.warning('No filter provided, returning all vectors')
return self.get(collection_name)
try:
@@ -443,17 +411,13 @@ class S3VectorClient(VectorDBBase):
all_vectors_result = self.get(collection_name)
if not all_vectors_result or not all_vectors_result.ids:
log.warning("No vectors found in collection")
log.warning('No vectors found in collection')
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
# Extract the lists from the result
all_ids = all_vectors_result.ids[0] if all_vectors_result.ids else []
all_documents = (
all_vectors_result.documents[0] if all_vectors_result.documents else []
)
all_metadatas = (
all_vectors_result.metadatas[0] if all_vectors_result.metadatas else []
)
all_documents = all_vectors_result.documents[0] if all_vectors_result.documents else []
all_metadatas = all_vectors_result.metadatas[0] if all_vectors_result.metadatas else []
# Apply client-side filtering
filtered_ids = []
@@ -472,9 +436,7 @@ class S3VectorClient(VectorDBBase):
if limit and len(filtered_ids) >= limit:
break
log.info(
f"Filter applied: {len(filtered_ids)} vectors match out of {len(all_ids)} total"
)
log.info(f'Filter applied: {len(filtered_ids)} vectors match out of {len(all_ids)} total')
# Return GetResult format
if filtered_ids:
@@ -489,15 +451,13 @@ class S3VectorClient(VectorDBBase):
except Exception as e:
log.error(f"Error querying collection '{collection_name}': {str(e)}")
# Handle specific AWS exceptions
if hasattr(e, "response") and "Error" in e.response:
error_code = e.response["Error"]["Code"]
if error_code == "NotFoundException":
if hasattr(e, 'response') and 'Error' in e.response:
error_code = e.response['Error']['Code']
if error_code == 'NotFoundException':
log.warning(f"Collection '{collection_name}' not found")
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
elif error_code == "AccessDeniedException":
log.error(
f"Access denied for collection '{collection_name}'. Check permissions."
)
elif error_code == 'AccessDeniedException':
log.error(f"Access denied for collection '{collection_name}'. Check permissions.")
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
raise
@@ -524,47 +484,43 @@ class S3VectorClient(VectorDBBase):
while True:
# Prepare request parameters
request_params = {
"vectorBucketName": self.bucket_name,
"indexName": collection_name,
"returnData": False, # Don't include vector data (not needed for get)
"returnMetadata": True, # Include metadata
"maxResults": 500, # Use reasonable page size
'vectorBucketName': self.bucket_name,
'indexName': collection_name,
'returnData': False, # Don't include vector data (not needed for get)
'returnMetadata': True, # Include metadata
'maxResults': 500, # Use reasonable page size
}
if next_token:
request_params["nextToken"] = next_token
request_params['nextToken'] = next_token
# Call S3 Vector API
response = self.client.list_vectors(**request_params)
# Process vectors in this page
vectors = response.get("vectors", [])
vectors = response.get('vectors', [])
for vector in vectors:
vector_id = vector.get("key")
vector_data = vector.get("data", {})
vector_metadata = vector.get("metadata", {})
vector_id = vector.get('key')
vector_data = vector.get('data', {})
vector_metadata = vector.get('metadata', {})
# Extract the actual vector array
vector_array = vector_data.get("float32", [])
vector_array = vector_data.get('float32', [])
# For documents, we try to extract text from metadata or use the vector ID
document_text = ""
document_text = ''
if isinstance(vector_metadata, dict):
# Get the text field first (highest priority)
document_text = vector_metadata.get("text")
document_text = vector_metadata.get('text')
if not document_text:
# Fallback to other possible text fields
document_text = (
vector_metadata.get("content")
or vector_metadata.get("document")
or vector_id
vector_metadata.get('content') or vector_metadata.get('document') or vector_id
)
# Log the actual content for debugging
log.debug(
f"Document text preview (first 200 chars): {str(document_text)[:200]}"
)
log.debug(f'Document text preview (first 200 chars): {str(document_text)[:200]}')
else:
document_text = vector_id
@@ -573,37 +529,29 @@ class S3VectorClient(VectorDBBase):
all_metadatas.append(vector_metadata)
# Check if there are more pages
next_token = response.get("nextToken")
next_token = response.get('nextToken')
if not next_token:
break
log.info(
f"Retrieved {len(all_ids)} vectors from collection '{collection_name}'"
)
log.info(f"Retrieved {len(all_ids)} vectors from collection '{collection_name}'")
# Return in GetResult format
# The Open WebUI GetResult expects lists of lists, so we wrap each list
if all_ids:
return GetResult(
ids=[all_ids], documents=[all_documents], metadatas=[all_metadatas]
)
return GetResult(ids=[all_ids], documents=[all_documents], metadatas=[all_metadatas])
else:
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
except Exception as e:
log.error(
f"Error retrieving vectors from collection '{collection_name}': {str(e)}"
)
log.error(f"Error retrieving vectors from collection '{collection_name}': {str(e)}")
# Handle specific AWS exceptions
if hasattr(e, "response") and "Error" in e.response:
error_code = e.response["Error"]["Code"]
if error_code == "NotFoundException":
if hasattr(e, 'response') and 'Error' in e.response:
error_code = e.response['Error']['Code']
if error_code == 'NotFoundException':
log.warning(f"Collection '{collection_name}' not found")
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
elif error_code == "AccessDeniedException":
log.error(
f"Access denied for collection '{collection_name}'. Check permissions."
)
elif error_code == 'AccessDeniedException':
log.error(f"Access denied for collection '{collection_name}'. Check permissions.")
return GetResult(ids=[[]], documents=[[]], metadatas=[[]])
raise
@@ -618,20 +566,16 @@ class S3VectorClient(VectorDBBase):
"""
if not self.has_collection(collection_name):
log.warning(
f"Collection '{collection_name}' does not exist, nothing to delete"
)
log.warning(f"Collection '{collection_name}' does not exist, nothing to delete")
return
# Check if this is a knowledge collection (not file-specific)
is_knowledge_collection = not collection_name.startswith("file-")
is_knowledge_collection = not collection_name.startswith('file-')
try:
if ids:
# Delete by specific vector IDs/keys
log.info(
f"Deleting {len(ids)} vectors by IDs from collection '{collection_name}'"
)
log.info(f"Deleting {len(ids)} vectors by IDs from collection '{collection_name}'")
self.client.delete_vectors(
vectorBucketName=self.bucket_name,
indexName=collection_name,
@@ -641,15 +585,13 @@ class S3VectorClient(VectorDBBase):
elif filter:
# Handle filter-based deletion
log.info(
f"Deleting vectors by filter from collection '{collection_name}': {filter}"
)
log.info(f"Deleting vectors by filter from collection '{collection_name}': {filter}")
# If this is a knowledge collection and we have a file_id filter,
# also clean up the corresponding file-specific collection
if is_knowledge_collection and "file_id" in filter:
file_id = filter["file_id"]
file_collection_name = f"file-{file_id}"
if is_knowledge_collection and 'file_id' in filter:
file_id = filter['file_id']
file_collection_name = f'file-{file_id}'
if self.has_collection(file_collection_name):
log.info(
f"Found related file-specific collection '{file_collection_name}', deleting it to prevent duplicates"
@@ -661,9 +603,7 @@ class S3VectorClient(VectorDBBase):
query_result = self.query(collection_name, filter)
if query_result and query_result.ids and query_result.ids[0]:
matching_ids = query_result.ids[0]
log.info(
f"Found {len(matching_ids)} vectors matching filter, deleting them"
)
log.info(f'Found {len(matching_ids)} vectors matching filter, deleting them')
# Delete the matching vectors by ID
self.client.delete_vectors(
@@ -671,17 +611,13 @@ class S3VectorClient(VectorDBBase):
indexName=collection_name,
keys=matching_ids,
)
log.info(
f"Deleted {len(matching_ids)} vectors from index '{collection_name}' using filter"
)
log.info(f"Deleted {len(matching_ids)} vectors from index '{collection_name}' using filter")
else:
log.warning("No vectors found matching the filter criteria")
log.warning('No vectors found matching the filter criteria')
else:
log.warning("No IDs or filter provided for deletion")
log.warning('No IDs or filter provided for deletion')
except Exception as e:
log.error(
f"Error deleting vectors from collection '{collection_name}': {e}"
)
log.error(f"Error deleting vectors from collection '{collection_name}': {e}")
raise
def reset(self) -> None:
@@ -690,36 +626,32 @@ class S3VectorClient(VectorDBBase):
"""
try:
log.warning(
"Reset called - this will delete all vector indexes in the S3 bucket"
)
log.warning('Reset called - this will delete all vector indexes in the S3 bucket')
# List all indexes
response = self.client.list_indexes(vectorBucketName=self.bucket_name)
indexes = response.get("indexes", [])
indexes = response.get('indexes', [])
if not indexes:
log.warning("No indexes found to delete")
log.warning('No indexes found to delete')
return
# Delete all indexes
deleted_count = 0
for index in indexes:
index_name = index.get("indexName")
index_name = index.get('indexName')
if index_name:
try:
self.client.delete_index(
vectorBucketName=self.bucket_name, indexName=index_name
)
self.client.delete_index(vectorBucketName=self.bucket_name, indexName=index_name)
deleted_count += 1
log.info(f"Deleted index: {index_name}")
log.info(f'Deleted index: {index_name}')
except Exception as e:
log.error(f"Error deleting index '{index_name}': {e}")
log.info(f"Reset completed: deleted {deleted_count} indexes")
log.info(f'Reset completed: deleted {deleted_count} indexes')
except Exception as e:
log.error(f"Error during reset: {e}")
log.error(f'Error during reset: {e}')
raise
def _matches_filter(self, metadata: Dict[str, Any], filter: Dict[str, Any]) -> bool:
@@ -732,15 +664,15 @@ class S3VectorClient(VectorDBBase):
# Check each filter condition
for key, expected_value in filter.items():
# Handle special operators
if key.startswith("$"):
if key == "$and":
if key.startswith('$'):
if key == '$and':
# All conditions must match
if not isinstance(expected_value, list):
continue
for condition in expected_value:
if not self._matches_filter(metadata, condition):
return False
elif key == "$or":
elif key == '$or':
# At least one condition must match
if not isinstance(expected_value, list):
continue
@@ -760,22 +692,19 @@ class S3VectorClient(VectorDBBase):
if isinstance(expected_value, dict):
# Handle comparison operators
for op, op_value in expected_value.items():
if op == "$eq":
if op == '$eq':
if actual_value != op_value:
return False
elif op == "$ne":
elif op == '$ne':
if actual_value == op_value:
return False
elif op == "$in":
if (
not isinstance(op_value, list)
or actual_value not in op_value
):
elif op == '$in':
if not isinstance(op_value, list) or actual_value not in op_value:
return False
elif op == "$nin":
elif op == '$nin':
if isinstance(op_value, list) and actual_value in op_value:
return False
elif op == "$exists":
elif op == '$exists':
if bool(op_value) != (key in metadata):
return False
# Add more operators as needed
@@ -60,47 +60,43 @@ class WeaviateClient(VectorDBBase):
try:
# Build connection parameters
connection_params = {
"http_host": WEAVIATE_HTTP_HOST,
"http_port": WEAVIATE_HTTP_PORT,
"http_secure": WEAVIATE_HTTP_SECURE,
"grpc_host": WEAVIATE_GRPC_HOST,
"grpc_port": WEAVIATE_GRPC_PORT,
"grpc_secure": WEAVIATE_GRPC_SECURE,
"skip_init_checks": WEAVIATE_SKIP_INIT_CHECKS,
'http_host': WEAVIATE_HTTP_HOST,
'http_port': WEAVIATE_HTTP_PORT,
'http_secure': WEAVIATE_HTTP_SECURE,
'grpc_host': WEAVIATE_GRPC_HOST,
'grpc_port': WEAVIATE_GRPC_PORT,
'grpc_secure': WEAVIATE_GRPC_SECURE,
'skip_init_checks': WEAVIATE_SKIP_INIT_CHECKS,
}
# Only add auth_credentials if WEAVIATE_API_KEY exists and is not empty
if WEAVIATE_API_KEY:
connection_params["auth_credentials"] = (
weaviate.classes.init.Auth.api_key(WEAVIATE_API_KEY)
)
connection_params['auth_credentials'] = weaviate.classes.init.Auth.api_key(WEAVIATE_API_KEY)
self.client = weaviate.connect_to_custom(**connection_params)
self.client.connect()
except Exception as e:
raise ConnectionError(f"Failed to connect to Weaviate: {e}") from e
raise ConnectionError(f'Failed to connect to Weaviate: {e}') from e
def _sanitize_collection_name(self, collection_name: str) -> str:
"""Sanitize collection name to be a valid Weaviate class name."""
if not isinstance(collection_name, str) or not collection_name.strip():
raise ValueError("Collection name must be a non-empty string")
raise ValueError('Collection name must be a non-empty string')
# Requirements for a valid Weaviate class name:
# The collection name must begin with a capital letter.
# The name can only contain letters, numbers, and the underscore (_) character. Spaces are not allowed.
# Replace hyphens with underscores and keep only alphanumeric characters
name = re.sub(r"[^a-zA-Z0-9_]", "", collection_name.replace("-", "_"))
name = name.strip("_")
name = re.sub(r'[^a-zA-Z0-9_]', '', collection_name.replace('-', '_'))
name = name.strip('_')
if not name:
raise ValueError(
"Could not sanitize collection name to be a valid Weaviate class name"
)
raise ValueError('Could not sanitize collection name to be a valid Weaviate class name')
# Ensure it starts with a letter and is capitalized
if not name[0].isalpha():
name = "C" + name
name = 'C' + name
return name[0].upper() + name[1:]
@@ -118,9 +114,7 @@ class WeaviateClient(VectorDBBase):
name=collection_name,
vector_config=weaviate.classes.config.Configure.Vectors.self_provided(),
properties=[
weaviate.classes.config.Property(
name="text", data_type=weaviate.classes.config.DataType.TEXT
),
weaviate.classes.config.Property(name='text', data_type=weaviate.classes.config.DataType.TEXT),
],
)
@@ -133,19 +127,15 @@ class WeaviateClient(VectorDBBase):
with collection.batch.fixed_size(batch_size=100) as batch:
for item in items:
item_uuid = str(uuid.uuid4()) if not item["id"] else str(item["id"])
item_uuid = str(uuid.uuid4()) if not item['id'] else str(item['id'])
properties = {"text": item["text"]}
if item["metadata"]:
clean_metadata = _convert_uuids_to_strings(
process_metadata(item["metadata"])
)
clean_metadata.pop("text", None)
properties = {'text': item['text']}
if item['metadata']:
clean_metadata = _convert_uuids_to_strings(process_metadata(item['metadata']))
clean_metadata.pop('text', None)
properties.update(clean_metadata)
batch.add_object(
properties=properties, uuid=item_uuid, vector=item["vector"]
)
batch.add_object(properties=properties, uuid=item_uuid, vector=item['vector'])
def upsert(self, collection_name: str, items: List[VectorItem]) -> None:
sane_collection_name = self._sanitize_collection_name(collection_name)
@@ -156,19 +146,15 @@ class WeaviateClient(VectorDBBase):
with collection.batch.fixed_size(batch_size=100) as batch:
for item in items:
item_uuid = str(item["id"]) if item["id"] else None
item_uuid = str(item['id']) if item['id'] else None
properties = {"text": item["text"]}
if item["metadata"]:
clean_metadata = _convert_uuids_to_strings(
process_metadata(item["metadata"])
)
clean_metadata.pop("text", None)
properties = {'text': item['text']}
if item['metadata']:
clean_metadata = _convert_uuids_to_strings(process_metadata(item['metadata']))
clean_metadata.pop('text', None)
properties.update(clean_metadata)
batch.add_object(
properties=properties, uuid=item_uuid, vector=item["vector"]
)
batch.add_object(properties=properties, uuid=item_uuid, vector=item['vector'])
def search(
self,
@@ -205,16 +191,12 @@ class WeaviateClient(VectorDBBase):
for obj in response.objects:
properties = dict(obj.properties) if obj.properties else {}
documents.append(properties.pop("text", ""))
documents.append(properties.pop('text', ''))
metadatas.append(_convert_uuids_to_strings(properties))
# Weaviate has cosine distance, 2 (worst) -> 0 (best). Re-ordering to 0 -> 1
raw_distances = [
(
obj.metadata.distance
if obj.metadata and obj.metadata.distance
else 2.0
)
(obj.metadata.distance if obj.metadata and obj.metadata.distance else 2.0)
for obj in response.objects
]
distances = [(2 - dist) / 2 for dist in raw_distances]
@@ -231,16 +213,14 @@ class WeaviateClient(VectorDBBase):
return SearchResult(
**{
"ids": result_ids,
"documents": result_documents,
"metadatas": result_metadatas,
"distances": result_distances,
'ids': result_ids,
'documents': result_documents,
'metadatas': result_metadatas,
'distances': result_distances,
}
)
def query(
self, collection_name: str, filter: Dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
sane_collection_name = self._sanitize_collection_name(collection_name)
if not self.client.collections.exists(sane_collection_name):
return None
@@ -250,21 +230,15 @@ class WeaviateClient(VectorDBBase):
weaviate_filter = None
if filter:
for key, value in filter.items():
prop_filter = weaviate.classes.query.Filter.by_property(name=key).equal(
value
)
prop_filter = weaviate.classes.query.Filter.by_property(name=key).equal(value)
weaviate_filter = (
prop_filter
if weaviate_filter is None
else weaviate.classes.query.Filter.all_of(
[weaviate_filter, prop_filter]
)
else weaviate.classes.query.Filter.all_of([weaviate_filter, prop_filter])
)
try:
response = collection.query.fetch_objects(
filters=weaviate_filter, limit=limit
)
response = collection.query.fetch_objects(filters=weaviate_filter, limit=limit)
ids = [str(obj.uuid) for obj in response.objects]
documents = []
@@ -272,14 +246,14 @@ class WeaviateClient(VectorDBBase):
for obj in response.objects:
properties = dict(obj.properties) if obj.properties else {}
documents.append(properties.pop("text", ""))
documents.append(properties.pop('text', ''))
metadatas.append(_convert_uuids_to_strings(properties))
return GetResult(
**{
"ids": [ids],
"documents": [documents],
"metadatas": [metadatas],
'ids': [ids],
'documents': [documents],
'metadatas': [metadatas],
}
)
except Exception:
@@ -297,7 +271,7 @@ class WeaviateClient(VectorDBBase):
for item in collection.iterator():
ids.append(str(item.uuid))
properties = dict(item.properties) if item.properties else {}
documents.append(properties.pop("text", ""))
documents.append(properties.pop('text', ''))
metadatas.append(_convert_uuids_to_strings(properties))
if not ids:
@@ -305,9 +279,9 @@ class WeaviateClient(VectorDBBase):
return GetResult(
**{
"ids": [ids],
"documents": [documents],
"metadatas": [metadatas],
'ids': [ids],
'documents': [documents],
'metadatas': [metadatas],
}
)
except Exception:
@@ -332,15 +306,11 @@ class WeaviateClient(VectorDBBase):
elif filter:
weaviate_filter = None
for key, value in filter.items():
prop_filter = weaviate.classes.query.Filter.by_property(
name=key
).equal(value)
prop_filter = weaviate.classes.query.Filter.by_property(name=key).equal(value)
weaviate_filter = (
prop_filter
if weaviate_filter is None
else weaviate.classes.query.Filter.all_of(
[weaviate_filter, prop_filter]
)
else weaviate.classes.query.Filter.all_of([weaviate_filter, prop_filter])
)
if weaviate_filter:
@@ -8,7 +8,6 @@ from open_webui.config import (
class Vector:
@staticmethod
def get_vector(vector_type: str) -> VectorDBBase:
"""
@@ -82,7 +81,7 @@ class Vector:
return WeaviateClient()
case _:
raise ValueError(f"Unsupported vector type: {vector_type}")
raise ValueError(f'Unsupported vector type: {vector_type}')
VECTOR_DB_CLIENT = Vector.get_vector(VECTOR_DB)
+1 -3
View File
@@ -63,9 +63,7 @@ class VectorDBBase(ABC):
pass
@abstractmethod
def query(
self, collection_name: str, filter: Dict, limit: Optional[int] = None
) -> Optional[GetResult]:
def query(self, collection_name: str, filter: Dict, limit: Optional[int] = None) -> Optional[GetResult]:
"""Query vectors from a collection using metadata filter."""
pass
+12 -12
View File
@@ -2,15 +2,15 @@ from enum import StrEnum
class VectorType(StrEnum):
MILVUS = "milvus"
MARIADB_VECTOR = "mariadb-vector"
QDRANT = "qdrant"
CHROMA = "chroma"
PINECONE = "pinecone"
ELASTICSEARCH = "elasticsearch"
OPENSEARCH = "opensearch"
PGVECTOR = "pgvector"
ORACLE23AI = "oracle23ai"
S3VECTOR = "s3vector"
WEAVIATE = "weaviate"
OPENGAUSS = "opengauss"
MILVUS = 'milvus'
MARIADB_VECTOR = 'mariadb-vector'
QDRANT = 'qdrant'
CHROMA = 'chroma'
PINECONE = 'pinecone'
ELASTICSEARCH = 'elasticsearch'
OPENSEARCH = 'opensearch'
PGVECTOR = 'pgvector'
ORACLE23AI = 'oracle23ai'
S3VECTOR = 's3vector'
WEAVIATE = 'weaviate'
OPENGAUSS = 'opengauss'
+2 -4
View File
@@ -1,13 +1,11 @@
from datetime import datetime
KEYS_TO_EXCLUDE = ["content", "pages", "tables", "paragraphs", "sections", "figures"]
KEYS_TO_EXCLUDE = ['content', 'pages', 'tables', 'paragraphs', 'sections', 'figures']
def filter_metadata(metadata: dict[str, any]) -> dict[str, any]:
# Removes large/redundant fields from metadata dict.
metadata = {
key: value for key, value in metadata.items() if key not in KEYS_TO_EXCLUDE
}
metadata = {key: value for key, value in metadata.items() if key not in KEYS_TO_EXCLUDE}
return metadata