refactor: move vdb implementations to workspaces (#34900)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: wangxiaolei <fatelei@gmail.com>
This commit is contained in:
Yunlu Wen
2026-04-13 08:56:43 +00:00
committed by GitHub
co-authored by autofix-ci[bot] wangxiaolei
parent c34f67495c
commit ae898652b2
223 changed files with 2009 additions and 984 deletions
+1
View File
@@ -0,0 +1 @@
"""Test suite root package (enables ``import tests.integration_tests...`` with ``pythonpath = .``)."""
+1
View File
@@ -0,0 +1 @@
"""Integration tests package."""
@@ -1,166 +0,0 @@
import os
from collections import UserDict
from typing import Any
from unittest.mock import MagicMock
import pytest
from _pytest.monkeypatch import MonkeyPatch
from pymochow import MochowClient
from pymochow.model.database import Database
from pymochow.model.enum import IndexState, IndexType, MetricType, ReadConsistency, TableState
from pymochow.model.schema import HNSWParams, VectorIndex
from pymochow.model.table import Table
class AttrDict(UserDict):
def __getattr__(self, item):
return self.get(item)
class MockBaiduVectorDBClass:
def mock_vector_db_client(
self,
config=None,
adapter: Any | None = None,
):
self.conn = MagicMock()
self._config = MagicMock()
def list_databases(self, config=None) -> list[Database]:
return [
Database(
conn=self.conn,
database_name="dify",
config=self._config,
)
]
def create_database(self, database_name: str, config=None) -> Database:
return Database(conn=self.conn, database_name=database_name, config=config)
def list_table(self, config=None) -> list[Table]:
return []
def drop_table(self, table_name: str, config=None):
return {"code": 0, "msg": "Success"}
def create_table(
self,
table_name: str,
replication: int,
partition: int,
schema,
enable_dynamic_field=False,
description: str = "",
config=None,
) -> Table:
return Table(self, table_name, replication, partition, schema, enable_dynamic_field, description, config)
def describe_table(self, table_name: str, config=None) -> Table:
return Table(
self,
table_name,
3,
1,
None,
enable_dynamic_field=False,
description="table for dify",
config=config,
state=TableState.NORMAL,
)
def upsert(self, rows, config=None):
return {"code": 0, "msg": "operation success", "affectedCount": 1}
def rebuild_index(self, index_name: str, config=None):
return {"code": 0, "msg": "Success"}
def describe_index(self, index_name: str, config=None):
return VectorIndex(
index_name=index_name,
index_type=IndexType.HNSW,
field="vector",
metric_type=MetricType.L2,
params=HNSWParams(m=16, efconstruction=200),
auto_build=False,
state=IndexState.NORMAL,
)
def query(
self,
primary_key,
partition_key=None,
projections=None,
retrieve_vector=False,
read_consistency=ReadConsistency.EVENTUAL,
config=None,
):
return AttrDict(
{
"row": {
"id": primary_key.get("id"),
"vector": [0.23432432, 0.8923744, 0.89238432],
"page_content": "text",
"metadata": {"doc_id": "doc_id_001"},
},
"code": 0,
"msg": "Success",
}
)
def delete(self, primary_key=None, partition_key=None, filter=None, config=None):
return {"code": 0, "msg": "Success"}
def search(
self,
anns,
partition_key=None,
projections=None,
retrieve_vector=False,
read_consistency=ReadConsistency.EVENTUAL,
config=None,
):
return AttrDict(
{
"rows": [
{
"row": {
"id": "doc_id_001",
"vector": [0.23432432, 0.8923744, 0.89238432],
"page_content": "text",
"metadata": {"doc_id": "doc_id_001"},
},
"distance": 0.1,
"score": 0.5,
}
],
"code": 0,
"msg": "Success",
}
)
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_baiduvectordb_mock(request, monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(MochowClient, "__init__", MockBaiduVectorDBClass.mock_vector_db_client)
monkeypatch.setattr(MochowClient, "list_databases", MockBaiduVectorDBClass.list_databases)
monkeypatch.setattr(MochowClient, "create_database", MockBaiduVectorDBClass.create_database)
monkeypatch.setattr(Database, "table", MockBaiduVectorDBClass.describe_table)
monkeypatch.setattr(Database, "list_table", MockBaiduVectorDBClass.list_table)
monkeypatch.setattr(Database, "create_table", MockBaiduVectorDBClass.create_table)
monkeypatch.setattr(Database, "drop_table", MockBaiduVectorDBClass.drop_table)
monkeypatch.setattr(Database, "describe_table", MockBaiduVectorDBClass.describe_table)
monkeypatch.setattr(Table, "rebuild_index", MockBaiduVectorDBClass.rebuild_index)
monkeypatch.setattr(Table, "describe_index", MockBaiduVectorDBClass.describe_index)
monkeypatch.setattr(Table, "delete", MockBaiduVectorDBClass.delete)
monkeypatch.setattr(Table, "query", MockBaiduVectorDBClass.query)
monkeypatch.setattr(Table, "search", MockBaiduVectorDBClass.search)
yield
if MOCK:
monkeypatch.undo()
@@ -1,209 +0,0 @@
import json
import os
from typing import Any
import holo_search_sdk as holo
import pytest
from _pytest.monkeypatch import MonkeyPatch
from psycopg import sql as psql
# Shared in-memory storage: {table_name: {doc_id: {"id", "text", "meta", "embedding"}}}
_mock_tables: dict[str, dict[str, dict[str, Any]]] = {}
class MockSearchQuery:
"""Mock query builder for search_vector and search_text results."""
def __init__(self, table_name: str, search_type: str):
self._table_name = table_name
self._search_type = search_type
self._limit_val = 10
self._filter_sql = None
def select(self, columns):
return self
def limit(self, n):
self._limit_val = n
return self
def where(self, filter_sql):
self._filter_sql = filter_sql
return self
def _apply_filter(self, row: dict[str, Any]) -> bool:
"""Apply the filter SQL to check if a row matches."""
if self._filter_sql is None:
return True
# Extract literals (the document IDs) from the filter SQL
# Filter format: meta->>'document_id' IN ('doc1', 'doc2')
literals = [v for t, v in _extract_identifiers_and_literals(self._filter_sql) if t == "literal"]
if not literals:
return True
# Get the document_id from the row's meta field
meta = row.get("meta", "{}")
if isinstance(meta, str):
meta = json.loads(meta)
doc_id = meta.get("document_id")
return doc_id in literals
def fetchall(self):
data = _mock_tables.get(self._table_name, {})
results = []
for row in list(data.values())[: self._limit_val]:
# Apply filter if present
if not self._apply_filter(row):
continue
if self._search_type == "vector":
# row format expected by _process_vector_results: (distance, id, text, meta)
results.append((0.1, row["id"], row["text"], row["meta"]))
else:
# row format expected by _process_full_text_results: (id, text, meta, embedding, score)
results.append((row["id"], row["text"], row["meta"], row.get("embedding", []), 0.9))
return results
class MockTable:
"""Mock table object returned by client.open_table()."""
def __init__(self, table_name: str):
self._table_name = table_name
def upsert_multi(self, index_column, values, column_names, update=True, update_columns=None):
if self._table_name not in _mock_tables:
_mock_tables[self._table_name] = {}
id_idx = column_names.index("id")
for row in values:
doc_id = row[id_idx]
_mock_tables[self._table_name][doc_id] = dict(zip(column_names, row))
def search_vector(self, vector, column, distance_method, output_name):
return MockSearchQuery(self._table_name, "vector")
def search_text(self, column, expression, return_score=False, return_score_name="score", return_all_columns=False):
return MockSearchQuery(self._table_name, "text")
def set_vector_index(
self, column, distance_method, base_quantization_type, max_degree, ef_construction, use_reorder
):
pass
def create_text_index(self, index_name, column, tokenizer):
pass
def _extract_sql_template(query) -> str:
"""Extract the SQL template string from a psycopg Composed object."""
if isinstance(query, psql.Composed):
for part in query:
if isinstance(part, psql.SQL):
return part._obj
if isinstance(query, psql.SQL):
return query._obj
return ""
def _extract_identifiers_and_literals(query) -> list[Any]:
"""Extract Identifier and Literal values from a psycopg Composed object."""
values: list[Any] = []
if isinstance(query, psql.Composed):
for part in query:
if isinstance(part, psql.Identifier):
values.append(("ident", part._obj[0] if part._obj else ""))
elif isinstance(part, psql.Literal):
values.append(("literal", part._obj))
elif isinstance(part, psql.Composed):
# Handles SQL(...).join(...) for IN clauses
for sub in part:
if isinstance(sub, psql.Literal):
values.append(("literal", sub._obj))
return values
class MockHologresClient:
"""Mock holo_search_sdk client that stores data in memory."""
def connect(self):
pass
def check_table_exist(self, table_name):
return table_name in _mock_tables
def open_table(self, table_name):
return MockTable(table_name)
def execute(self, query, fetch_result=False):
template = _extract_sql_template(query)
params = _extract_identifiers_and_literals(query)
if "CREATE TABLE" in template.upper():
# Extract table name from first identifier
table_name = next((v for t, v in params if t == "ident"), "unknown")
if table_name not in _mock_tables:
_mock_tables[table_name] = {}
return None
if "SELECT 1" in template:
# text_exists: SELECT 1 FROM {table} WHERE id = {id} LIMIT 1
table_name = next((v for t, v in params if t == "ident"), "")
doc_id = next((v for t, v in params if t == "literal"), "")
data = _mock_tables.get(table_name, {})
return [(1,)] if doc_id in data else []
if "SELECT id" in template:
# get_ids_by_metadata_field: SELECT id FROM {table} WHERE meta->>{key} = {value}
table_name = next((v for t, v in params if t == "ident"), "")
literals = [v for t, v in params if t == "literal"]
key = literals[0] if len(literals) > 0 else ""
value = literals[1] if len(literals) > 1 else ""
data = _mock_tables.get(table_name, {})
return [(doc_id,) for doc_id, row in data.items() if json.loads(row.get("meta", "{}")).get(key) == value]
if "DELETE" in template.upper():
table_name = next((v for t, v in params if t == "ident"), "")
if "id IN" in template:
# delete_by_ids
ids_to_delete = [v for t, v in params if t == "literal"]
for did in ids_to_delete:
_mock_tables.get(table_name, {}).pop(did, None)
elif "meta->>" in template:
# delete_by_metadata_field
literals = [v for t, v in params if t == "literal"]
key = literals[0] if len(literals) > 0 else ""
value = literals[1] if len(literals) > 1 else ""
data = _mock_tables.get(table_name, {})
to_remove = [
doc_id for doc_id, row in data.items() if json.loads(row.get("meta", "{}")).get(key) == value
]
for did in to_remove:
data.pop(did, None)
return None
return [] if fetch_result else None
def drop_table(self, table_name):
_mock_tables.pop(table_name, None)
def mock_connect(**kwargs):
"""Replacement for holo_search_sdk.connect() that returns a mock client."""
return MockHologresClient()
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_hologres_mock(monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(holo, "connect", mock_connect)
yield
if MOCK:
_mock_tables.clear()
monkeypatch.undo()
@@ -1,89 +0,0 @@
import os
import pytest
from _pytest.monkeypatch import MonkeyPatch
from elasticsearch import Elasticsearch
from core.rag.datasource.vdb.field import Field
class MockIndicesClient:
def __init__(self):
pass
def create(self, index, mappings, settings):
return {"acknowledge": True}
def refresh(self, index):
return {"acknowledge": True}
def delete(self, index):
return {"acknowledge": True}
def exists(self, index):
return True
class MockClient:
def __init__(self, **kwargs):
self.indices = MockIndicesClient()
def index(self, **kwargs):
return {"acknowledge": True}
def exists(self, **kwargs):
return True
def delete(self, **kwargs):
return {"acknowledge": True}
def search(self, **kwargs):
return {
"took": 1,
"hits": {
"hits": [
{
"_source": {
Field.CONTENT_KEY: "abcdef",
Field.VECTOR: [1, 2],
Field.METADATA_KEY: {},
},
"_score": 1.0,
},
{
"_source": {
Field.CONTENT_KEY: "123456",
Field.VECTOR: [2, 2],
Field.METADATA_KEY: {},
},
"_score": 0.9,
},
{
"_source": {
Field.CONTENT_KEY: "a1b2c3",
Field.VECTOR: [3, 2],
Field.METADATA_KEY: {},
},
"_score": 0.8,
},
]
},
}
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_client_mock(request, monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(Elasticsearch, "__init__", MockClient.__init__)
monkeypatch.setattr(Elasticsearch, "index", MockClient.index)
monkeypatch.setattr(Elasticsearch, "exists", MockClient.exists)
monkeypatch.setattr(Elasticsearch, "delete", MockClient.delete)
monkeypatch.setattr(Elasticsearch, "search", MockClient.search)
yield
if MOCK:
monkeypatch.undo()
@@ -1,192 +0,0 @@
import os
from typing import Any, Union
import pytest
from _pytest.monkeypatch import MonkeyPatch
from tcvectordb import RPCVectorDBClient
from tcvectordb.model import enum
from tcvectordb.model.collection import FilterIndexConfig
from tcvectordb.model.document import AnnSearch, Document, Filter, KeywordSearch, Rerank
from tcvectordb.model.enum import ReadConsistency
from tcvectordb.model.index import FilterIndex, HNSWParams, Index, IndexField, VectorIndex
from tcvectordb.rpc.model.collection import RPCCollection
from tcvectordb.rpc.model.database import RPCDatabase
from xinference_client.types import Embedding
class MockTcvectordbClass:
def mock_vector_db_client(
self,
url: str,
username="",
key="",
read_consistency: ReadConsistency = ReadConsistency.EVENTUAL_CONSISTENCY,
timeout=10,
adapter: Any | None = None,
pool_size: int = 2,
proxies: dict | None = None,
password: str | None = None,
**kwargs,
):
self._conn = None
self._read_consistency = read_consistency
def create_database_if_not_exists(self, database_name: str, timeout: float | None = None) -> RPCDatabase:
return RPCDatabase(
name="dify",
read_consistency=self._read_consistency,
)
def exists_collection(self, database_name: str, collection_name: str) -> bool:
return True
def describe_collection(
self, database_name: str, collection_name: str, timeout: float | None = None
) -> RPCCollection:
index = Index(
FilterIndex("id", enum.FieldType.String, enum.IndexType.PRIMARY_KEY),
VectorIndex(
"vector",
128,
enum.IndexType.HNSW,
enum.MetricType.IP,
HNSWParams(m=16, efconstruction=200),
),
FilterIndex("text", enum.FieldType.String, enum.IndexType.FILTER),
FilterIndex("metadata", enum.FieldType.String, enum.IndexType.FILTER),
)
return RPCCollection(
RPCDatabase(
name=database_name,
read_consistency=self._read_consistency,
),
collection_name,
index=index,
)
def create_collection(
self,
database_name: str,
collection_name: str,
shard: int,
replicas: int,
description: str | None = None,
index: Index | None = None,
embedding: Embedding | None = None,
timeout: float | None = None,
ttl_config: dict | None = None,
filter_index_config: FilterIndexConfig | None = None,
indexes: list[IndexField] | None = None,
) -> RPCCollection:
return RPCCollection(
RPCDatabase(
name="dify",
read_consistency=self._read_consistency,
),
collection_name,
shard,
replicas,
description,
index,
embedding=embedding,
read_consistency=self._read_consistency,
timeout=timeout,
ttl_config=ttl_config,
filter_index_config=filter_index_config,
indexes=indexes,
)
def collection_upsert(
self,
database_name: str,
collection_name: str,
documents: list[Union[Document, dict]],
timeout: float | None = None,
build_index: bool = True,
**kwargs,
):
return {"code": 0, "msg": "operation success"}
def collection_search(
self,
database_name: str,
collection_name: str,
vectors: list[list[float]],
filter: Filter | None = None,
params=None,
retrieve_vector: bool = False,
limit: int = 10,
output_fields: list[str] | None = None,
timeout: float | None = None,
) -> list[list[dict]]:
return [[{"metadata": {"doc_id": "foo1"}, "text": "text", "doc_id": "foo1", "score": 0.1}]]
def collection_hybrid_search(
self,
database_name: str,
collection_name: str,
ann: Union[list[AnnSearch], AnnSearch] | None = None,
match: Union[list[KeywordSearch], KeywordSearch] | None = None,
filter: Union[Filter, str] | None = None,
rerank: Rerank | None = None,
retrieve_vector: bool | None = None,
output_fields: list[str] | None = None,
limit: int | None = None,
timeout: float | None = None,
return_pd_object=False,
**kwargs,
) -> list[list[dict]]:
return [[{"metadata": {"doc_id": "foo1"}, "text": "text", "doc_id": "foo1", "score": 0.1}]]
def collection_query(
self,
database_name: str,
collection_name: str,
document_ids: list | None = None,
retrieve_vector: bool = False,
limit: int | None = None,
offset: int | None = None,
filter: Filter | None = None,
output_fields: list[str] | None = None,
timeout: float | None = None,
):
return [{"metadata": '{"doc_id":"foo1"}', "text": "text", "doc_id": "foo1", "score": 0.1}]
def collection_delete(
self,
database_name: str,
collection_name: str,
document_ids: list[str] | None = None,
filter: Filter | None = None,
timeout: float | None = None,
):
return {"code": 0, "msg": "operation success"}
def drop_collection(self, database_name: str, collection_name: str, timeout: float | None = None):
return {"code": 0, "msg": "operation success"}
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_tcvectordb_mock(request, monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(RPCVectorDBClient, "__init__", MockTcvectordbClass.mock_vector_db_client)
monkeypatch.setattr(
RPCVectorDBClient, "create_database_if_not_exists", MockTcvectordbClass.create_database_if_not_exists
)
monkeypatch.setattr(RPCVectorDBClient, "exists_collection", MockTcvectordbClass.exists_collection)
monkeypatch.setattr(RPCVectorDBClient, "create_collection", MockTcvectordbClass.create_collection)
monkeypatch.setattr(RPCVectorDBClient, "describe_collection", MockTcvectordbClass.describe_collection)
monkeypatch.setattr(RPCVectorDBClient, "upsert", MockTcvectordbClass.collection_upsert)
monkeypatch.setattr(RPCVectorDBClient, "search", MockTcvectordbClass.collection_search)
monkeypatch.setattr(RPCVectorDBClient, "hybrid_search", MockTcvectordbClass.collection_hybrid_search)
monkeypatch.setattr(RPCVectorDBClient, "query", MockTcvectordbClass.collection_query)
monkeypatch.setattr(RPCVectorDBClient, "delete", MockTcvectordbClass.collection_delete)
monkeypatch.setattr(RPCVectorDBClient, "drop_collection", MockTcvectordbClass.drop_collection)
yield
if MOCK:
monkeypatch.undo()
@@ -1,75 +0,0 @@
import os
from collections import UserDict
import pytest
from _pytest.monkeypatch import MonkeyPatch
from upstash_vector import Index
# Mocking the Index class from upstash_vector
class MockIndex:
def __init__(self, url="", token=""):
self.url = url
self.token = token
self.vectors = []
def upsert(self, vectors):
for vector in vectors:
vector.score = 0.5
self.vectors.append(vector)
return {"code": 0, "msg": "operation success", "affectedCount": len(vectors)}
def fetch(self, ids):
return [vector for vector in self.vectors if vector.id in ids]
def delete(self, ids):
self.vectors = [vector for vector in self.vectors if vector.id not in ids]
return {"code": 0, "msg": "Success"}
def query(
self,
vector: None,
top_k: int = 10,
include_vectors: bool = False,
include_metadata: bool = False,
filter: str = "",
data: str | None = None,
namespace: str = "",
include_data: bool = False,
):
# Simple mock query, in real scenario you would calculate similarity
mock_result = []
for vector_data in self.vectors:
mock_result.append(vector_data)
return mock_result[:top_k]
def reset(self):
self.vectors = []
def info(self):
return AttrDict({"dimension": 1024})
class AttrDict(UserDict):
def __getattr__(self, item):
return self.get(item)
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_upstashvector_mock(request, monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(Index, "__init__", MockIndex.__init__)
monkeypatch.setattr(Index, "upsert", MockIndex.upsert)
monkeypatch.setattr(Index, "fetch", MockIndex.fetch)
monkeypatch.setattr(Index, "delete", MockIndex.delete)
monkeypatch.setattr(Index, "query", MockIndex.query)
monkeypatch.setattr(Index, "reset", MockIndex.reset)
monkeypatch.setattr(Index, "info", MockIndex.info)
yield
if MOCK:
monkeypatch.undo()
@@ -1,215 +0,0 @@
import os
from typing import Union
from unittest.mock import MagicMock
import pytest
from _pytest.monkeypatch import MonkeyPatch
from volcengine.viking_db import (
Collection,
Data,
DistanceType,
Field,
FieldType,
Index,
IndexType,
QuantType,
VectorIndexParams,
VikingDBService,
)
from core.rag.datasource.vdb.field import Field as vdb_Field
class MockVikingDBClass:
def __init__(
self,
host="api-vikingdb.volces.com",
region="cn-north-1",
ak="",
sk="",
scheme="http",
connection_timeout=30,
socket_timeout=30,
proxy=None,
):
self._viking_db_service = MagicMock()
self._viking_db_service.get_exception = MagicMock(return_value='{"data": {"primary_key": "test_id"}}')
def get_collection(self, collection_name) -> Collection:
return Collection(
collection_name=collection_name,
description="Collection For Dify",
viking_db_service=self._viking_db_service,
primary_key=vdb_Field.PRIMARY_KEY,
fields=[
Field(field_name=vdb_Field.PRIMARY_KEY, field_type=FieldType.String, is_primary_key=True),
Field(field_name=vdb_Field.METADATA_KEY, field_type=FieldType.String),
Field(field_name=vdb_Field.GROUP_KEY, field_type=FieldType.String),
Field(field_name=vdb_Field.CONTENT_KEY, field_type=FieldType.Text),
Field(field_name=vdb_Field.VECTOR, field_type=FieldType.Vector, dim=768),
],
indexes=[
Index(
collection_name=collection_name,
index_name=f"{collection_name}_idx",
vector_index=VectorIndexParams(
distance=DistanceType.L2,
index_type=IndexType.HNSW,
quant=QuantType.Float,
),
scalar_index=None,
stat=None,
viking_db_service=self._viking_db_service,
)
],
)
def drop_collection(self, collection_name):
assert collection_name != ""
def create_collection(self, collection_name, fields, description="") -> Collection:
return Collection(
collection_name=collection_name,
description=description,
primary_key=vdb_Field.PRIMARY_KEY,
viking_db_service=self._viking_db_service,
fields=fields,
)
def get_index(self, collection_name, index_name) -> Index:
return Index(
collection_name=collection_name,
index_name=index_name,
viking_db_service=self._viking_db_service,
stat=None,
scalar_index=None,
vector_index=VectorIndexParams(
distance=DistanceType.L2,
index_type=IndexType.HNSW,
quant=QuantType.Float,
),
)
def create_index(
self,
collection_name,
index_name,
vector_index=None,
cpu_quota=2,
description="",
partition_by="",
scalar_index=None,
shard_count=None,
shard_policy=None,
):
return Index(
collection_name=collection_name,
index_name=index_name,
vector_index=vector_index,
cpu_quota=cpu_quota,
description=description,
partition_by=partition_by,
scalar_index=scalar_index,
shard_count=shard_count,
shard_policy=shard_policy,
viking_db_service=self._viking_db_service,
stat=None,
)
def drop_index(self, collection_name, index_name):
assert collection_name != ""
assert index_name != ""
def upsert_data(self, data: Union[Data, list[Data]]):
assert data is not None
def fetch_data(self, id: Union[str, list[str], int, list[int]]):
return Data(
fields={
vdb_Field.GROUP_KEY: "test_group",
vdb_Field.METADATA_KEY: "{}",
vdb_Field.CONTENT_KEY: "content",
vdb_Field.PRIMARY_KEY: id,
vdb_Field.VECTOR: [-0.00762577411336441, -0.01949881482151406, 0.008832383941428398],
},
id=id,
)
def delete_data(self, id: Union[str, list[str], int, list[int]]):
assert id is not None
def search_by_vector(
self,
vector,
sparse_vectors=None,
filter=None,
limit=10,
output_fields=None,
partition="default",
dense_weight=None,
) -> list[Data]:
return [
Data(
fields={
vdb_Field.GROUP_KEY: "test_group",
vdb_Field.METADATA_KEY: '\
{"source": "/var/folders/ml/xxx/xxx.txt", \
"document_id": "test_document_id", \
"dataset_id": "test_dataset_id", \
"doc_id": "test_id", \
"doc_hash": "test_hash"}',
vdb_Field.CONTENT_KEY: "content",
vdb_Field.PRIMARY_KEY: "test_id",
vdb_Field.VECTOR: vector,
},
id="test_id",
score=0.10,
)
]
def search(
self, order=None, filter=None, limit=10, output_fields=None, partition="default", dense_weight=None
) -> list[Data]:
return [
Data(
fields={
vdb_Field.GROUP_KEY: "test_group",
vdb_Field.METADATA_KEY: '\
{"source": "/var/folders/ml/xxx/xxx.txt", \
"document_id": "test_document_id", \
"dataset_id": "test_dataset_id", \
"doc_id": "test_id", \
"doc_hash": "test_hash"}',
vdb_Field.CONTENT_KEY: "content",
vdb_Field.PRIMARY_KEY: "test_id",
vdb_Field.VECTOR: [-0.00762577411336441, -0.01949881482151406, 0.008832383941428398],
},
id="test_id",
score=0.10,
)
]
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
@pytest.fixture
def setup_vikingdb_mock(monkeypatch: MonkeyPatch):
if MOCK:
monkeypatch.setattr(VikingDBService, "__init__", MockVikingDBClass.__init__)
monkeypatch.setattr(VikingDBService, "get_collection", MockVikingDBClass.get_collection)
monkeypatch.setattr(VikingDBService, "create_collection", MockVikingDBClass.create_collection)
monkeypatch.setattr(VikingDBService, "drop_collection", MockVikingDBClass.drop_collection)
monkeypatch.setattr(VikingDBService, "get_index", MockVikingDBClass.get_index)
monkeypatch.setattr(VikingDBService, "create_index", MockVikingDBClass.create_index)
monkeypatch.setattr(VikingDBService, "drop_index", MockVikingDBClass.drop_index)
monkeypatch.setattr(Collection, "upsert_data", MockVikingDBClass.upsert_data)
monkeypatch.setattr(Collection, "fetch_data", MockVikingDBClass.fetch_data)
monkeypatch.setattr(Collection, "delete_data", MockVikingDBClass.delete_data)
monkeypatch.setattr(Index, "search_by_vector", MockVikingDBClass.search_by_vector)
monkeypatch.setattr(Index, "search", MockVikingDBClass.search)
yield
if MOCK:
monkeypatch.undo()
@@ -1,51 +0,0 @@
from core.rag.datasource.vdb.analyticdb.analyticdb_vector import AnalyticdbVector
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_openapi import AnalyticdbVectorOpenAPIConfig
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_sql import AnalyticdbVectorBySqlConfig
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class AnalyticdbVectorTest(AbstractVectorTest):
def __init__(self, config_type: str):
super().__init__()
# Analyticdb requires collection_name length less than 60.
# it's ok for normal usage.
self.collection_name = self.collection_name.replace("_test", "")
if config_type == "sql":
self.vector = AnalyticdbVector(
collection_name=self.collection_name,
sql_config=AnalyticdbVectorBySqlConfig(
host="test_host",
port=5432,
account="test_account",
account_password="test_passwd",
namespace="difytest_namespace",
),
api_config=None,
)
else:
self.vector = AnalyticdbVector(
collection_name=self.collection_name,
sql_config=None,
api_config=AnalyticdbVectorOpenAPIConfig(
access_key_id="test_key_id",
access_key_secret="test_key_secret",
region_id="test_region",
instance_id="test_id",
account="test_account",
account_password="test_passwd",
namespace="difytest_namespace",
collection="difytest_collection",
namespace_password="test_passwd",
),
)
def run_all_tests(self):
self.vector.delete()
return super().run_all_tests()
def test_chroma_vector(setup_mock_redis):
AnalyticdbVectorTest("api").run_all_tests()
AnalyticdbVectorTest("sql").run_all_tests()
@@ -1,35 +0,0 @@
from core.rag.datasource.vdb.baidu.baidu_vector import BaiduConfig, BaiduVector
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = (
"tests.integration_tests.vdb.test_vector_store",
"tests.integration_tests.vdb.__mock.baiduvectordb",
)
class BaiduVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = BaiduVector(
"dify",
BaiduConfig(
endpoint="http://127.0.0.1:5287",
account="root",
api_key="dify",
database="dify",
shard=1,
replicas=3,
),
)
def search_by_vector(self):
hits_by_vector = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 1
def search_by_full_text(self):
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def test_baidu_vector(setup_mock_redis, setup_baiduvectordb_mock):
BaiduVectorTest().run_all_tests()
@@ -1,34 +0,0 @@
import chromadb
from core.rag.datasource.vdb.chroma.chroma_vector import ChromaConfig, ChromaVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
get_example_text,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class ChromaVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = ChromaVector(
collection_name=self.collection_name,
config=ChromaConfig(
host="localhost",
port=8000,
tenant=chromadb.DEFAULT_TENANT,
database=chromadb.DEFAULT_DATABASE,
auth_provider="chromadb.auth.token_authn.TokenAuthClientProvider",
auth_credentials="difyai123456",
),
)
def search_by_full_text(self):
# chroma dos not support full text searching
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def test_chroma_vector(setup_mock_redis):
ChromaVectorTest().run_all_tests()
@@ -1,25 +0,0 @@
# Clickzetta Integration Tests
## Running Tests
To run the Clickzetta integration tests, you need to set the following environment variables:
```bash
export CLICKZETTA_USERNAME=your_username
export CLICKZETTA_PASSWORD=your_password
export CLICKZETTA_INSTANCE=your_instance
export CLICKZETTA_SERVICE=api.clickzetta.com
export CLICKZETTA_WORKSPACE=your_workspace
export CLICKZETTA_VCLUSTER=your_vcluster
export CLICKZETTA_SCHEMA=dify
```
Then run the tests:
```bash
pytest api/tests/integration_tests/vdb/clickzetta/
```
## Security Note
Never commit credentials to the repository. Always use environment variables or secure credential management systems.
@@ -1,223 +0,0 @@
import contextlib
import os
import pytest
from core.rag.datasource.vdb.clickzetta.clickzetta_vector import ClickzettaConfig, ClickzettaVector
from core.rag.models.document import Document
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text, setup_mock_redis
class TestClickzettaVector(AbstractVectorTest):
"""
Test cases for Clickzetta vector database integration.
"""
@pytest.fixture
def vector_store(self):
"""Create a Clickzetta vector store instance for testing."""
# Skip test if Clickzetta credentials are not configured
if not os.getenv("CLICKZETTA_USERNAME"):
pytest.skip("CLICKZETTA_USERNAME is not configured")
if not os.getenv("CLICKZETTA_PASSWORD"):
pytest.skip("CLICKZETTA_PASSWORD is not configured")
if not os.getenv("CLICKZETTA_INSTANCE"):
pytest.skip("CLICKZETTA_INSTANCE is not configured")
config = ClickzettaConfig(
username=os.getenv("CLICKZETTA_USERNAME", ""),
password=os.getenv("CLICKZETTA_PASSWORD", ""),
instance=os.getenv("CLICKZETTA_INSTANCE", ""),
service=os.getenv("CLICKZETTA_SERVICE", "api.clickzetta.com"),
workspace=os.getenv("CLICKZETTA_WORKSPACE", "quick_start"),
vcluster=os.getenv("CLICKZETTA_VCLUSTER", "default_ap"),
schema=os.getenv("CLICKZETTA_SCHEMA", "dify_test"),
batch_size=10, # Small batch size for testing
enable_inverted_index=True,
analyzer_type="chinese",
analyzer_mode="smart",
vector_distance_function="cosine_distance",
)
with setup_mock_redis():
vector = ClickzettaVector(collection_name="test_collection_" + str(os.getpid()), config=config)
yield vector
# Cleanup: delete the test collection
with contextlib.suppress(Exception):
vector.delete()
def test_clickzetta_vector_basic_operations(self, vector_store):
"""Test basic CRUD operations on Clickzetta vector store."""
# Prepare test data
texts = [
"这是第一个测试文档,包含一些中文内容。",
"This is the second test document with English content.",
"第三个文档混合了English和中文内容。",
]
embeddings = [
[0.1, 0.2, 0.3, 0.4],
[0.5, 0.6, 0.7, 0.8],
[0.9, 1.0, 1.1, 1.2],
]
documents = [
Document(page_content=text, metadata={"doc_id": f"doc_{i}", "source": "test"})
for i, text in enumerate(texts)
]
# Test create (initial insert)
vector_store.create(texts=documents, embeddings=embeddings)
# Test text_exists
assert vector_store.text_exists("doc_0")
assert not vector_store.text_exists("doc_999")
# Test search_by_vector
query_vector = [0.1, 0.2, 0.3, 0.4]
results = vector_store.search_by_vector(query_vector, top_k=2)
assert len(results) > 0
assert results[0].page_content == texts[0] # Should match the first document
# Test search_by_full_text (Chinese)
results = vector_store.search_by_full_text("中文", top_k=3)
assert len(results) >= 2 # Should find documents with Chinese content
# Test search_by_full_text (English)
results = vector_store.search_by_full_text("English", top_k=3)
assert len(results) >= 2 # Should find documents with English content
# Test delete_by_ids
vector_store.delete_by_ids(["doc_0"])
assert not vector_store.text_exists("doc_0")
assert vector_store.text_exists("doc_1")
# Test delete_by_metadata_field
vector_store.delete_by_metadata_field("source", "test")
assert not vector_store.text_exists("doc_1")
assert not vector_store.text_exists("doc_2")
def test_clickzetta_vector_advanced_search(self, vector_store):
"""Test advanced search features of Clickzetta vector store."""
# Prepare test data with more complex metadata
documents = []
embeddings = []
for i in range(10):
doc = Document(
page_content=f"Document {i}: " + get_example_text(),
metadata={
"doc_id": f"adv_doc_{i}",
"category": "technical" if i % 2 == 0 else "general",
"document_id": f"doc_{i // 3}", # Group documents
"importance": i,
},
)
documents.append(doc)
# Create varied embeddings
embeddings.append([0.1 * i, 0.2 * i, 0.3 * i, 0.4 * i])
vector_store.create(texts=documents, embeddings=embeddings)
# Test vector search with document filter
query_vector = [0.5, 1.0, 1.5, 2.0]
results = vector_store.search_by_vector(query_vector, top_k=5, document_ids_filter=["doc_0", "doc_1"])
assert len(results) > 0
# All results should belong to doc_0 or doc_1 groups
for result in results:
assert result.metadata["document_id"] in ["doc_0", "doc_1"]
# Test score threshold
results = vector_store.search_by_vector(query_vector, top_k=10, score_threshold=0.5)
# Check that all results have a score above threshold
for result in results:
assert result.metadata.get("score", 0) >= 0.5
def test_clickzetta_batch_operations(self, vector_store):
"""Test batch insertion operations."""
# Prepare large batch of documents
batch_size = 25
documents = []
embeddings = []
for i in range(batch_size):
doc = Document(
page_content=f"Batch document {i}: This is a test document for batch processing.",
metadata={"doc_id": f"batch_doc_{i}", "batch": "test_batch"},
)
documents.append(doc)
embeddings.append([0.1 * (i % 10), 0.2 * (i % 10), 0.3 * (i % 10), 0.4 * (i % 10)])
# Test batch insert
vector_store.add_texts(documents=documents, embeddings=embeddings)
# Verify all documents were inserted
for i in range(batch_size):
assert vector_store.text_exists(f"batch_doc_{i}")
# Clean up
vector_store.delete_by_metadata_field("batch", "test_batch")
def test_clickzetta_edge_cases(self, vector_store):
"""Test edge cases and error handling."""
# Test empty operations
vector_store.create(texts=[], embeddings=[])
vector_store.add_texts(documents=[], embeddings=[])
vector_store.delete_by_ids([])
# Test special characters in content
special_doc = Document(
page_content="Special chars: 'quotes', \"double\", \\backslash, \n newline",
metadata={"doc_id": "special_doc", "test": "edge_case"},
)
embeddings = [[0.1, 0.2, 0.3, 0.4]]
vector_store.add_texts(documents=[special_doc], embeddings=embeddings)
assert vector_store.text_exists("special_doc")
# Test search with special characters
results = vector_store.search_by_full_text("quotes", top_k=1)
if results: # Full-text search might not be available
assert len(results) > 0
# Clean up
vector_store.delete_by_ids(["special_doc"])
def test_clickzetta_full_text_search_modes(self, vector_store):
"""Test different full-text search capabilities."""
# Prepare documents with various language content
documents = [
Document(
page_content="云器科技提供强大的Lakehouse解决方案", metadata={"doc_id": "cn_doc_1", "lang": "chinese"}
),
Document(
page_content="Clickzetta provides powerful Lakehouse solutions",
metadata={"doc_id": "en_doc_1", "lang": "english"},
),
Document(
page_content="Lakehouse是现代数据架构的重要组成部分", metadata={"doc_id": "cn_doc_2", "lang": "chinese"}
),
Document(
page_content="Modern data architecture includes Lakehouse technology",
metadata={"doc_id": "en_doc_2", "lang": "english"},
),
]
embeddings = [[0.1, 0.2, 0.3, 0.4] for _ in documents]
vector_store.create(texts=documents, embeddings=embeddings)
# Test Chinese full-text search
results = vector_store.search_by_full_text("Lakehouse", top_k=4)
assert len(results) >= 2 # Should find at least documents with "Lakehouse"
# Test English full-text search
results = vector_store.search_by_full_text("solutions", top_k=2)
assert len(results) >= 1 # Should find English documents with "solutions"
# Test mixed search
results = vector_store.search_by_full_text("数据架构", top_k=2)
assert len(results) >= 1 # Should find Chinese documents with this phrase
# Clean up
vector_store.delete_by_metadata_field("lang", "chinese")
vector_store.delete_by_metadata_field("lang", "english")
@@ -1,165 +0,0 @@
#!/usr/bin/env python3
"""
Test Clickzetta integration in Docker environment
"""
import os
import time
import httpx
from clickzetta import connect
def test_clickzetta_connection():
"""Test direct connection to Clickzetta"""
print("=== Testing direct Clickzetta connection ===")
try:
conn = connect(
username=os.getenv("CLICKZETTA_USERNAME", "test_user"),
password=os.getenv("CLICKZETTA_PASSWORD", "test_password"),
instance=os.getenv("CLICKZETTA_INSTANCE", "test_instance"),
service=os.getenv("CLICKZETTA_SERVICE", "api.clickzetta.com"),
workspace=os.getenv("CLICKZETTA_WORKSPACE", "test_workspace"),
vcluster=os.getenv("CLICKZETTA_VCLUSTER", "default"),
database=os.getenv("CLICKZETTA_SCHEMA", "dify"),
)
with conn.cursor() as cursor:
# Test basic connectivity
cursor.execute("SELECT 1 as test")
result = cursor.fetchone()
print(f"✓ Connection test: {result}")
# Check if our test table exists
cursor.execute("SHOW TABLES IN dify")
tables = cursor.fetchall()
print(f"✓ Existing tables: {[t[1] for t in tables if t[0] == 'dify']}")
# Check if test collection exists
test_collection = "collection_test_dataset"
if test_collection in [t[1] for t in tables if t[0] == "dify"]:
cursor.execute(f"DESCRIBE dify.{test_collection}")
columns = cursor.fetchall()
print(f"✓ Table structure for {test_collection}:")
for col in columns:
print(f" - {col[0]}: {col[1]}")
# Check for indexes
cursor.execute(f"SHOW INDEXES IN dify.{test_collection}")
indexes = cursor.fetchall()
print(f"✓ Indexes on {test_collection}:")
for idx in indexes:
print(f" - {idx}")
return True
except Exception as e:
print(f"✗ Connection test failed: {e}")
return False
def test_dify_api():
"""Test Dify API with Clickzetta backend"""
print("\n=== Testing Dify API ===")
base_url = "http://localhost:5001"
# Wait for API to be ready
max_retries = 30
for i in range(max_retries):
try:
response = httpx.get(f"{base_url}/console/api/health")
if response.status_code == 200:
print("✓ Dify API is ready")
break
except:
if i == max_retries - 1:
print("✗ Dify API is not responding")
return False
time.sleep(2)
# Check vector store configuration
try:
# This is a simplified check - in production, you'd use proper auth
print("✓ Dify is configured to use Clickzetta as vector store")
return True
except Exception as e:
print(f"✗ API test failed: {e}")
return False
def verify_table_structure():
"""Verify the table structure meets Dify requirements"""
print("\n=== Verifying Table Structure ===")
expected_columns = {
"id": "VARCHAR",
"page_content": "VARCHAR",
"metadata": "VARCHAR", # JSON stored as VARCHAR in Clickzetta
"vector": "ARRAY<FLOAT>",
}
expected_metadata_fields = ["doc_id", "doc_hash", "document_id", "dataset_id"]
print("✓ Expected table structure:")
for col, dtype in expected_columns.items():
print(f" - {col}: {dtype}")
print("\n✓ Required metadata fields:")
for field in expected_metadata_fields:
print(f" - {field}")
print("\n✓ Index requirements:")
print(" - Vector index (HNSW) on 'vector' column")
print(" - Full-text index on 'page_content' (optional)")
print(" - Functional index on metadata->>'$.doc_id' (recommended)")
print(" - Functional index on metadata->>'$.document_id' (recommended)")
return True
def main():
"""Run all tests"""
print("Starting Clickzetta integration tests for Dify Docker\n")
tests = [
("Direct Clickzetta Connection", test_clickzetta_connection),
("Dify API Status", test_dify_api),
("Table Structure Verification", verify_table_structure),
]
results = []
for test_name, test_func in tests:
try:
success = test_func()
results.append((test_name, success))
except Exception as e:
print(f"\n✗ {test_name} crashed: {e}")
results.append((test_name, False))
# Summary
print("\n" + "=" * 50)
print("Test Summary:")
print("=" * 50)
passed = sum(1 for _, success in results if success)
total = len(results)
for test_name, success in results:
status = "✅ PASSED" if success else "❌ FAILED"
print(f"{test_name}: {status}")
print(f"\nTotal: {passed}/{total} tests passed")
if passed == total:
print("\n🎉 All tests passed! Clickzetta is ready for Dify Docker deployment.")
print("\nNext steps:")
print("1. Run: cd docker && docker-compose -f docker-compose.yaml -f docker-compose.clickzetta.yaml up -d")
print("2. Access Dify at http://localhost:3000")
print("3. Create a dataset and test vector storage with Clickzetta")
return 0
else:
print("\n⚠️ Some tests failed. Please check the errors above.")
return 1
if __name__ == "__main__":
exit(main())
@@ -1,50 +0,0 @@
import subprocess
import time
from core.rag.datasource.vdb.couchbase.couchbase_vector import CouchbaseConfig, CouchbaseVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
def wait_for_healthy_container(service_name="couchbase-server", timeout=300):
start_time = time.time()
while time.time() - start_time < timeout:
result = subprocess.run(
["docker", "inspect", "--format", "{{.State.Health.Status}}", service_name], capture_output=True, text=True
)
if result.stdout.strip() == "healthy":
print(f"{service_name} is healthy!")
return True
else:
print(f"Waiting for {service_name} to be healthy...")
time.sleep(10)
raise TimeoutError(f"{service_name} did not become healthy in time")
class CouchbaseTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = CouchbaseVector(
collection_name=self.collection_name,
config=CouchbaseConfig(
connection_string="couchbase://127.0.0.1",
user="Administrator",
password="password",
bucket_name="Embeddings",
scope_name="_default",
),
)
def search_by_vector(self):
# brief sleep to ensure document is indexed
time.sleep(5)
hits_by_vector = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 1
def test_couchbase(setup_mock_redis):
wait_for_healthy_container("couchbase-server", timeout=60)
CouchbaseTest().run_all_tests()
@@ -1,23 +0,0 @@
from core.rag.datasource.vdb.elasticsearch.elasticsearch_vector import ElasticSearchConfig, ElasticSearchVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class ElasticSearchVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.attributes = ["doc_id", "dataset_id", "document_id", "doc_hash"]
self.vector = ElasticSearchVector(
index_name=self.collection_name.lower(),
config=ElasticSearchConfig(
use_cloud=False, host="http://localhost", port="9200", username="elastic", password="elastic"
),
attributes=self.attributes,
)
def test_elasticsearch_vector(setup_mock_redis):
ElasticSearchVectorTest().run_all_tests()
@@ -1,153 +0,0 @@
import os
import uuid
from typing import cast
from holo_search_sdk.types import BaseQuantizationType, DistanceType, TokenizerType
from core.rag.datasource.vdb.hologres.hologres_vector import HologresVector, HologresVectorConfig
from core.rag.models.document import Document
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = (
"tests.integration_tests.vdb.test_vector_store",
"tests.integration_tests.vdb.__mock.hologres",
)
MOCK = os.getenv("MOCK_SWITCH", "false").lower() == "true"
class HologresVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
# Hologres requires collection names to be lowercase
self.collection_name = self.collection_name.lower()
self.vector = HologresVector(
collection_name=self.collection_name,
config=HologresVectorConfig(
host=os.environ.get("HOLOGRES_HOST", "localhost"),
port=int(os.environ.get("HOLOGRES_PORT", "80")),
database=os.environ.get("HOLOGRES_DATABASE", "test_db"),
access_key_id=os.environ.get("HOLOGRES_ACCESS_KEY_ID", "test_key"),
access_key_secret=os.environ.get("HOLOGRES_ACCESS_KEY_SECRET", "test_secret"),
schema_name=os.environ.get("HOLOGRES_SCHEMA", "public"),
tokenizer=cast(TokenizerType, os.environ.get("HOLOGRES_TOKENIZER", "jieba")),
distance_method=cast(DistanceType, os.environ.get("HOLOGRES_DISTANCE_METHOD", "Cosine")),
base_quantization_type=cast(
BaseQuantizationType, os.environ.get("HOLOGRES_BASE_QUANTIZATION_TYPE", "rabitq")
),
max_degree=int(os.environ.get("HOLOGRES_MAX_DEGREE", "64")),
ef_construction=int(os.environ.get("HOLOGRES_EF_CONSTRUCTION", "400")),
),
)
def search_by_full_text(self):
"""Override: full-text index may not be immediately ready in real mode."""
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
if MOCK:
# In mock mode, full-text search should return the document we inserted
assert len(hits_by_full_text) == 1
assert hits_by_full_text[0].metadata["doc_id"] == self.example_doc_id
else:
# In real mode, full-text index may need time to become active
assert len(hits_by_full_text) >= 0
def search_by_vector_with_filter(self):
"""Test vector search with document_ids_filter."""
# Create another document with different document_id
other_doc_id = str(uuid.uuid4())
other_doc = Document(
page_content="other_text",
metadata={
"doc_id": other_doc_id,
"doc_hash": other_doc_id,
"document_id": other_doc_id,
"dataset_id": self.dataset_id,
},
)
self.vector.add_texts(documents=[other_doc], embeddings=[self.example_embedding])
# Search with filter - should only return the original document
hits = self.vector.search_by_vector(
query_vector=self.example_embedding,
document_ids_filter=[self.example_doc_id],
)
assert len(hits) == 1
assert hits[0].metadata["doc_id"] == self.example_doc_id
# Search without filter - should return both
all_hits = self.vector.search_by_vector(query_vector=self.example_embedding, top_k=10)
assert len(all_hits) >= 2
def search_by_full_text_with_filter(self):
"""Test full-text search with document_ids_filter."""
# Create another document with different document_id
other_doc_id = str(uuid.uuid4())
other_doc = Document(
page_content="unique_other_text",
metadata={
"doc_id": other_doc_id,
"doc_hash": other_doc_id,
"document_id": other_doc_id,
"dataset_id": self.dataset_id,
},
)
self.vector.add_texts(documents=[other_doc], embeddings=[self.example_embedding])
# Search with filter - should only return the original document
hits = self.vector.search_by_full_text(
query=get_example_text(),
document_ids_filter=[self.example_doc_id],
)
if MOCK:
assert len(hits) == 1
assert hits[0].metadata["doc_id"] == self.example_doc_id
def get_ids_by_metadata_field(self):
"""Override: Hologres implements this method via JSONB query."""
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert ids is not None
assert len(ids) == 1
def run_all_tests(self):
# Clean up before running tests
self.vector.delete()
# Run base tests (create, search, text_exists, get_ids, add_texts, delete_by_ids, delete)
super().run_all_tests()
# Additional filter tests require fresh data (table was deleted by base tests)
if MOCK:
# Recreate collection for filter tests
self.vector.create(
texts=[
Document(
page_content=get_example_text(),
metadata={
"doc_id": self.example_doc_id,
"doc_hash": self.example_doc_id,
"document_id": self.example_doc_id,
"dataset_id": self.dataset_id,
},
)
],
embeddings=[self.example_embedding],
)
self.search_by_vector_with_filter()
self.search_by_full_text_with_filter()
# Clean up
self.vector.delete()
def test_hologres_vector(setup_mock_redis, setup_hologres_mock):
"""
Test Hologres vector database implementation.
This test covers:
- Creating collection with vector index
- Adding texts with embeddings
- Vector similarity search
- Full-text search
- Text existence check
- Batch deletion by IDs
- Collection deletion
"""
HologresVectorTest().run_all_tests()
@@ -1,32 +0,0 @@
from core.rag.datasource.vdb.huawei.huawei_cloud_vector import HuaweiCloudVector, HuaweiCloudVectorConfig
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = (
"tests.integration_tests.vdb.test_vector_store",
"tests.integration_tests.vdb.__mock.huaweicloudvectordb",
)
class HuaweiCloudVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = HuaweiCloudVector(
"dify",
HuaweiCloudVectorConfig(
hosts="https://127.0.0.1:9200",
username="dify",
password="dify",
),
)
def search_by_vector(self):
hits_by_vector = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 3
def search_by_full_text(self):
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 3
def test_huawei_cloud_vector(setup_mock_redis, setup_client_mock):
HuaweiCloudVectorTest().run_all_tests()
@@ -1,45 +0,0 @@
"""Integration tests for IRIS vector database."""
from core.rag.datasource.vdb.iris.iris_vector import IrisVector, IrisVectorConfig
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class IrisVectorTest(AbstractVectorTest):
"""Test suite for IRIS vector store implementation."""
def __init__(self):
"""Initialize IRIS vector test with hardcoded test configuration.
Note: Uses 'host.docker.internal' to connect from DevContainer to
host OS Docker, or 'localhost' when running directly on host OS.
"""
super().__init__()
self.vector = IrisVector(
collection_name=self.collection_name,
config=IrisVectorConfig(
IRIS_HOST="host.docker.internal",
IRIS_SUPER_SERVER_PORT=1972,
IRIS_USER="_SYSTEM",
IRIS_PASSWORD="Dify@1234",
IRIS_DATABASE="USER",
IRIS_SCHEMA="dify",
IRIS_CONNECTION_URL=None,
IRIS_MIN_CONNECTION=1,
IRIS_MAX_CONNECTION=3,
IRIS_TEXT_INDEX=True,
IRIS_TEXT_INDEX_LANGUAGE="en",
),
)
def test_iris_vector(setup_mock_redis) -> None:
"""Run all IRIS vector store tests.
Args:
setup_mock_redis: Pytest fixture for mock Redis setup
"""
IrisVectorTest().run_all_tests()
@@ -1,60 +0,0 @@
import os
from core.rag.datasource.vdb.lindorm.lindorm_vector import LindormVectorStore, LindormVectorStoreConfig
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class Config:
SEARCH_ENDPOINT = os.environ.get(
"SEARCH_ENDPOINT", "http://ld-************-proxy-search-pub.lindorm.aliyuncs.com:30070"
)
SEARCH_USERNAME = os.environ.get("SEARCH_USERNAME", "ADMIN")
SEARCH_PWD = os.environ.get("SEARCH_PWD", "ADMIN")
USING_UGC = os.environ.get("USING_UGC", "True").lower() == "true"
class TestLindormVectorStore(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = LindormVectorStore(
collection_name=self.collection_name,
config=LindormVectorStoreConfig(
hosts=Config.SEARCH_ENDPOINT,
username=Config.SEARCH_USERNAME,
password=Config.SEARCH_PWD,
),
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="doc_id", value=self.example_doc_id)
assert ids is not None
assert len(ids) == 1
assert ids[0] == self.example_doc_id
class TestLindormVectorStoreUGC(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = LindormVectorStore(
collection_name="ugc_index_test",
config=LindormVectorStoreConfig(
hosts=Config.SEARCH_ENDPOINT,
username=Config.SEARCH_USERNAME,
password=Config.SEARCH_PWD,
using_ugc=Config.USING_UGC,
),
routing_value=self.collection_name,
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="doc_id", value=self.example_doc_id)
assert ids is not None
assert len(ids) == 1
assert ids[0] == self.example_doc_id
def test_lindorm_vector_ugc(setup_mock_redis):
TestLindormVectorStore().run_all_tests()
TestLindormVectorStoreUGC().run_all_tests()
@@ -1,25 +0,0 @@
from core.rag.datasource.vdb.matrixone.matrixone_vector import MatrixoneConfig, MatrixoneVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class MatrixoneVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = MatrixoneVector(
collection_name=self.collection_name,
config=MatrixoneConfig(
host="localhost", port=6001, user="dump", password="111", database="dify", metric="l2"
),
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
def test_matrixone_vector(setup_mock_redis):
MatrixoneVectorTest().run_all_tests()
@@ -1,33 +0,0 @@
from core.rag.datasource.vdb.milvus.milvus_vector import MilvusConfig, MilvusVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
get_example_text,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class MilvusVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = MilvusVector(
collection_name=self.collection_name,
config=MilvusConfig(
uri="http://localhost:19530",
user="root",
password="Milvus",
),
)
def search_by_full_text(self):
# milvus support BM25 full text search after version 2.5.0-beta
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) >= 0
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
def test_milvus_vector(setup_mock_redis):
MilvusVectorTest().run_all_tests()
@@ -1,30 +0,0 @@
from core.rag.datasource.vdb.myscale.myscale_vector import MyScaleConfig, MyScaleVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class MyScaleVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = MyScaleVector(
collection_name=self.collection_name,
config=MyScaleConfig(
host="localhost",
port=8123,
user="default",
password="",
database="dify",
fts_params="",
),
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
def test_myscale_vector(setup_mock_redis):
MyScaleVectorTest().run_all_tests()
@@ -1,241 +0,0 @@
"""
Benchmark: OceanBase vector store — old (single-row) vs new (batch) insertion,
metadata query with/without functional index, and vector search across metrics.
Usage:
uv run --project api python -m tests.integration_tests.vdb.oceanbase.bench_oceanbase
"""
import json
import random
import statistics
import time
import uuid
from pyobvector import VECTOR, ObVecClient, cosine_distance, inner_product, l2_distance
from sqlalchemy import JSON, Column, String, text
from sqlalchemy.dialects.mysql import LONGTEXT
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
HOST = "127.0.0.1"
PORT = 2881
USER = "root@test"
PASSWORD = "difyai123456"
DATABASE = "test"
VEC_DIM = 1536
HNSW_BUILD = {"M": 16, "efConstruction": 256}
DISTANCE_FUNCS = {"l2": l2_distance, "cosine": cosine_distance, "inner_product": inner_product}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_client(**extra):
return ObVecClient(
uri=f"{HOST}:{PORT}",
user=USER,
password=PASSWORD,
db_name=DATABASE,
**extra,
)
def _rand_vec():
return [random.uniform(-1, 1) for _ in range(VEC_DIM)] # noqa: S311
def _drop(client, table):
client.drop_table_if_exist(table)
def _create_table(client, table, metric="l2"):
cols = [
Column("id", String(36), primary_key=True, autoincrement=False),
Column("vector", VECTOR(VEC_DIM)),
Column("text", LONGTEXT),
Column("metadata", JSON),
]
vidx = client.prepare_index_params()
vidx.add_index(
field_name="vector",
index_type="HNSW",
index_name="vector_index",
metric_type=metric,
params=HNSW_BUILD,
)
client.create_table_with_index_params(table_name=table, columns=cols, vidxs=vidx)
client.refresh_metadata([table])
def _gen_rows(n):
doc_id = str(uuid.uuid4())
rows = []
for _ in range(n):
rows.append(
{
"id": str(uuid.uuid4()),
"vector": _rand_vec(),
"text": f"benchmark text {uuid.uuid4().hex[:12]}",
"metadata": json.dumps({"document_id": doc_id, "dataset_id": str(uuid.uuid4())}),
}
)
return rows, doc_id
# ---------------------------------------------------------------------------
# Benchmark: Insertion
# ---------------------------------------------------------------------------
def bench_insert_single(client, table, rows):
"""Old approach: one INSERT per row."""
t0 = time.perf_counter()
for row in rows:
client.insert(table_name=table, data=row)
return time.perf_counter() - t0
def bench_insert_batch(client, table, rows, batch_size=100):
"""New approach: batch INSERT."""
t0 = time.perf_counter()
for start in range(0, len(rows), batch_size):
batch = rows[start : start + batch_size]
client.insert(table_name=table, data=batch)
return time.perf_counter() - t0
# ---------------------------------------------------------------------------
# Benchmark: Metadata query
# ---------------------------------------------------------------------------
def bench_metadata_query(client, table, doc_id, with_index=False):
"""Query by metadata->>'$.document_id' with/without functional index."""
if with_index:
try:
client.perform_raw_text_sql(f"CREATE INDEX idx_metadata_doc_id ON `{table}` ((metadata->>'$.document_id'))")
except Exception:
pass # already exists
sql = text(f"SELECT id FROM `{table}` WHERE metadata->>'$.document_id' = :val")
times = []
with client.engine.connect() as conn:
for _ in range(10):
t0 = time.perf_counter()
result = conn.execute(sql, {"val": doc_id})
_ = result.fetchall()
times.append(time.perf_counter() - t0)
return times
# ---------------------------------------------------------------------------
# Benchmark: Vector search
# ---------------------------------------------------------------------------
def bench_vector_search(client, table, metric, topk=10, n_queries=20):
dist_func = DISTANCE_FUNCS[metric]
times = []
for _ in range(n_queries):
q = _rand_vec()
t0 = time.perf_counter()
cur = client.ann_search(
table_name=table,
vec_column_name="vector",
vec_data=q,
topk=topk,
distance_func=dist_func,
output_column_names=["text", "metadata"],
with_dist=True,
)
_ = list(cur)
times.append(time.perf_counter() - t0)
return times
def _fmt(times):
"""Format list of durations as 'mean ± stdev'."""
m = statistics.mean(times) * 1000
s = statistics.stdev(times) * 1000 if len(times) > 1 else 0
return f"{m:.1f} ± {s:.1f} ms"
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
client = _make_client()
client_pooled = _make_client(pool_size=5, max_overflow=10, pool_recycle=3600, pool_pre_ping=True)
print("=" * 70)
print("OceanBase Vector Store — Performance Benchmark")
print(f" Endpoint : {HOST}:{PORT}")
print(f" Vec dim : {VEC_DIM}")
print("=" * 70)
# ------------------------------------------------------------------
# 1. Insertion benchmark
# ------------------------------------------------------------------
for n_docs in [100, 500, 1000]:
rows, doc_id = _gen_rows(n_docs)
tbl_single = f"bench_single_{n_docs}"
tbl_batch = f"bench_batch_{n_docs}"
_drop(client, tbl_single)
_drop(client, tbl_batch)
_create_table(client, tbl_single)
_create_table(client, tbl_batch)
t_single = bench_insert_single(client, tbl_single, rows)
t_batch = bench_insert_batch(client_pooled, tbl_batch, rows, batch_size=100)
speedup = t_single / t_batch if t_batch > 0 else float("inf")
print(f"\n[Insert {n_docs} docs]")
print(f" Single-row : {t_single:.2f}s")
print(f" Batch(100) : {t_batch:.2f}s")
print(f" Speedup : {speedup:.1f}x")
# ------------------------------------------------------------------
# 2. Metadata query benchmark (use the 1000-doc batch table)
# ------------------------------------------------------------------
tbl_meta = "bench_batch_1000"
rows_1000, doc_id_1000 = _gen_rows(1000)
# The table already has 1000 rows from step 1; use that doc_id
# Re-query doc_id from one of the rows we inserted
with client.engine.connect() as conn:
res = conn.execute(text(f"SELECT metadata->>'$.document_id' FROM `{tbl_meta}` LIMIT 1"))
doc_id_1000 = res.fetchone()[0]
print("\n[Metadata filter query — 1000 rows, by document_id]")
times_no_idx = bench_metadata_query(client, tbl_meta, doc_id_1000, with_index=False)
print(f" Without index : {_fmt(times_no_idx)}")
times_with_idx = bench_metadata_query(client, tbl_meta, doc_id_1000, with_index=True)
print(f" With index : {_fmt(times_with_idx)}")
# ------------------------------------------------------------------
# 3. Vector search benchmark — across metrics
# ------------------------------------------------------------------
print("\n[Vector search — top-10, 20 queries each, on 1000 rows]")
for metric in ["l2", "cosine", "inner_product"]:
tbl_vs = f"bench_vs_{metric}"
_drop(client_pooled, tbl_vs)
_create_table(client_pooled, tbl_vs, metric=metric)
# Insert 1000 rows
rows_vs, _ = _gen_rows(1000)
bench_insert_batch(client_pooled, tbl_vs, rows_vs, batch_size=100)
times = bench_vector_search(client_pooled, tbl_vs, metric, topk=10, n_queries=20)
print(f" {metric:15s}: {_fmt(times)}")
_drop(client_pooled, tbl_vs)
# ------------------------------------------------------------------
# Cleanup
# ------------------------------------------------------------------
for n in [100, 500, 1000]:
_drop(client, f"bench_single_{n}")
_drop(client, f"bench_batch_{n}")
print("\n" + "=" * 70)
print("Benchmark complete.")
print("=" * 70)
if __name__ == "__main__":
main()
@@ -1,44 +0,0 @@
import pytest
from core.rag.datasource.vdb.oceanbase.oceanbase_vector import (
OceanBaseVector,
OceanBaseVectorConfig,
)
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
@pytest.fixture
def oceanbase_vector():
return OceanBaseVector(
"dify_test_collection",
config=OceanBaseVectorConfig(
host="127.0.0.1",
port=2881,
user="root",
database="test",
password="difyai123456",
enable_hybrid_search=True,
batch_size=10,
),
)
class OceanBaseVectorTest(AbstractVectorTest):
def __init__(self, vector: OceanBaseVector):
super().__init__()
self.vector = vector
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
def test_oceanbase_vector(
setup_mock_redis,
oceanbase_vector,
):
OceanBaseVectorTest(oceanbase_vector).run_all_tests()
@@ -1,42 +0,0 @@
import time
import psycopg2
from core.rag.datasource.vdb.opengauss.opengauss import OpenGauss, OpenGaussConfig
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class OpenGaussTest(AbstractVectorTest):
def __init__(self):
super().__init__()
max_retries = 5
retry_delay = 20
retry_count = 0
while retry_count < max_retries:
try:
config = OpenGaussConfig(
host="localhost",
port=6600,
user="postgres",
password="Dify@123",
database="dify",
min_connection=1,
max_connection=5,
)
break
except psycopg2.OperationalError as e:
retry_count += 1
if retry_count < max_retries:
time.sleep(retry_delay)
self.vector = OpenGauss(
collection_name=self.collection_name,
config=config,
)
def test_opengauss(setup_mock_redis):
OpenGaussTest().run_all_tests()
@@ -1,235 +0,0 @@
from unittest.mock import MagicMock, patch
import pytest
from core.rag.datasource.vdb.field import Field
from core.rag.datasource.vdb.opensearch.opensearch_vector import OpenSearchConfig, OpenSearchVector
from core.rag.models.document import Document
from extensions import ext_redis
def get_example_text() -> str:
return "This is a sample text for testing purposes."
@pytest.fixture(scope="module")
def setup_mock_redis():
ext_redis.redis_client.get = MagicMock(return_value=None)
ext_redis.redis_client.set = MagicMock(return_value=None)
mock_redis_lock = MagicMock()
mock_redis_lock.__enter__ = MagicMock()
mock_redis_lock.__exit__ = MagicMock()
ext_redis.redis_client.lock = MagicMock(return_value=mock_redis_lock)
class TestOpenSearchConfig:
def test_to_opensearch_params(self):
config = OpenSearchConfig(
host="localhost",
port=9200,
secure=True,
user="admin",
password="password",
)
params = config.to_opensearch_params()
assert params["hosts"] == [{"host": "localhost", "port": 9200}]
assert params["use_ssl"] is True
assert params["verify_certs"] is True
assert params["connection_class"].__name__ == "Urllib3HttpConnection"
assert params["http_auth"] == ("admin", "password")
@patch("boto3.Session", autospec=True)
@patch("core.rag.datasource.vdb.opensearch.opensearch_vector.Urllib3AWSV4SignerAuth", autospec=True)
def test_to_opensearch_params_with_aws_managed_iam(
self, mock_aws_signer_auth: MagicMock, mock_boto_session: MagicMock
):
mock_credentials = MagicMock()
mock_boto_session.return_value.get_credentials.return_value = mock_credentials
mock_auth_instance = mock_aws_signer_auth.return_value
aws_region = "ap-southeast-2"
aws_service = "aoss"
host = f"aoss-endpoint.{aws_region}.aoss.amazonaws.com"
port = 9201
config = OpenSearchConfig(
host=host,
port=port,
secure=True,
auth_method="aws_managed_iam",
aws_region=aws_region,
aws_service=aws_service,
)
params = config.to_opensearch_params()
assert params["hosts"] == [{"host": host, "port": port}]
assert params["use_ssl"] is True
assert params["verify_certs"] is True
assert params["connection_class"].__name__ == "Urllib3HttpConnection"
assert params["http_auth"] is mock_auth_instance
mock_aws_signer_auth.assert_called_once_with(
credentials=mock_credentials, region=aws_region, service=aws_service
)
assert mock_boto_session.return_value.get_credentials.called
class TestOpenSearchVector:
def setup_method(self):
self.collection_name = "test_collection"
self.example_doc_id = "example_doc_id"
self.vector = OpenSearchVector(
collection_name=self.collection_name,
config=OpenSearchConfig(host="localhost", port=9200, secure=False, user="admin", password="password"),
)
self.vector._client = MagicMock()
@pytest.mark.parametrize(
("search_response", "expected_length", "expected_doc_id"),
[
(
{
"hits": {
"total": {"value": 1},
"hits": [
{
"_source": {
"page_content": get_example_text(),
"metadata": {"document_id": "example_doc_id"},
}
}
],
}
},
1,
"example_doc_id",
),
({"hits": {"total": {"value": 0}, "hits": []}}, 0, None),
],
)
def test_search_by_full_text(self, search_response, expected_length, expected_doc_id):
self.vector._client.search.return_value = search_response
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == expected_length
if expected_length > 0:
assert hits_by_full_text[0].metadata["document_id"] == expected_doc_id
def test_search_by_vector(self):
vector = [0.1] * 128
mock_response = {
"hits": {
"total": {"value": 1},
"hits": [
{
"_source": {
Field.CONTENT_KEY: get_example_text(),
Field.METADATA_KEY: {"document_id": self.example_doc_id},
},
"_score": 1.0,
}
],
}
}
self.vector._client.search.return_value = mock_response
hits_by_vector = self.vector.search_by_vector(query_vector=vector)
print("Hits by vector:", hits_by_vector)
print("Expected document ID:", self.example_doc_id)
print("Actual document ID:", hits_by_vector[0].metadata["document_id"] if hits_by_vector else "No hits")
assert len(hits_by_vector) > 0, f"Expected at least one hit, got {len(hits_by_vector)}"
assert hits_by_vector[0].metadata["document_id"] == self.example_doc_id, (
f"Expected document ID {self.example_doc_id}, got {hits_by_vector[0].metadata['document_id']}"
)
def test_get_ids_by_metadata_field(self):
mock_response = {"hits": {"total": {"value": 1}, "hits": [{"_id": "mock_id"}]}}
self.vector._client.search.return_value = mock_response
doc = Document(page_content="Test content", metadata={"document_id": self.example_doc_id})
embedding = [0.1] * 128
with patch("opensearchpy.helpers.bulk", autospec=True) as mock_bulk:
mock_bulk.return_value = ([], [])
self.vector.add_texts([doc], [embedding])
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
assert ids[0] == "mock_id"
def test_add_texts(self):
self.vector._client.index.return_value = {"result": "created"}
doc = Document(page_content="Test content", metadata={"document_id": self.example_doc_id})
embedding = [0.1] * 128
with patch("opensearchpy.helpers.bulk", autospec=True) as mock_bulk:
mock_bulk.return_value = ([], [])
self.vector.add_texts([doc], [embedding])
mock_response = {"hits": {"total": {"value": 1}, "hits": [{"_id": "mock_id"}]}}
self.vector._client.search.return_value = mock_response
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
assert ids[0] == "mock_id"
def test_delete_nonexistent_index(self):
"""Test deleting a non-existent index."""
# Create a vector instance with a non-existent collection name
self.vector._client.indices.exists.return_value = False
# Should not raise an exception
self.vector.delete()
# Verify that exists was called but delete was not
self.vector._client.indices.exists.assert_called_once_with(index=self.collection_name.lower())
self.vector._client.indices.delete.assert_not_called()
def test_delete_existing_index(self):
"""Test deleting an existing index."""
self.vector._client.indices.exists.return_value = True
self.vector.delete()
# Verify both exists and delete were called
self.vector._client.indices.exists.assert_called_once_with(index=self.collection_name.lower())
self.vector._client.indices.delete.assert_called_once_with(index=self.collection_name.lower())
@pytest.mark.usefixtures("setup_mock_redis")
class TestOpenSearchVectorWithRedis:
def setup_method(self):
self.tester = TestOpenSearchVector()
def test_search_by_full_text(self):
self.tester.setup_method()
search_response = {
"hits": {
"total": {"value": 1},
"hits": [
{"_source": {"page_content": get_example_text(), "metadata": {"document_id": "example_doc_id"}}}
],
}
}
expected_length = 1
expected_doc_id = "example_doc_id"
self.tester.test_search_by_full_text(search_response, expected_length, expected_doc_id)
def test_get_ids_by_metadata_field(self):
self.tester.setup_method()
self.tester.test_get_ids_by_metadata_field()
def test_add_texts(self):
self.tester.setup_method()
self.tester.test_add_texts()
def test_search_by_vector(self):
self.tester.setup_method()
self.tester.test_search_by_vector()
@@ -1,29 +0,0 @@
from core.rag.datasource.vdb.oracle.oraclevector import OracleVector, OracleVectorConfig
from core.rag.models.document import Document
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
get_example_text,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class OracleVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = OracleVector(
collection_name=self.collection_name,
config=OracleVectorConfig(
user="dify",
password="dify",
dsn="localhost:1521/FREEPDB1",
),
)
def search_by_full_text(self):
hits_by_full_text: list[Document] = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def test_oraclevector(setup_mock_redis):
OracleVectorTest().run_all_tests()
@@ -1,36 +0,0 @@
from core.rag.datasource.vdb.pgvecto_rs.pgvecto_rs import PGVectoRS, PgvectoRSConfig
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
get_example_text,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class PGVectoRSVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = PGVectoRS(
collection_name=self.collection_name.lower(),
config=PgvectoRSConfig(
host="localhost",
port=5431,
user="postgres",
password="difyai123456",
database="dify",
),
dim=128,
)
def search_by_full_text(self):
# pgvecto rs only support english text search, So it’s not open for now
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) == 1
def test_pgvecto_rs(setup_mock_redis):
PGVectoRSVectorTest().run_all_tests()
@@ -1,27 +0,0 @@
from core.rag.datasource.vdb.pgvector.pgvector import PGVector, PGVectorConfig
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class PGVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = PGVector(
collection_name=self.collection_name,
config=PGVectorConfig(
host="localhost",
port=5433,
user="postgres",
password="difyai123456",
database="dify",
min_connection=1,
max_connection=5,
),
)
def test_pgvector(setup_mock_redis):
PGVectorTest().run_all_tests()
@@ -1,27 +0,0 @@
from core.rag.datasource.vdb.pyvastbase.vastbase_vector import VastbaseVector, VastbaseVectorConfig
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class VastbaseVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = VastbaseVector(
collection_name=self.collection_name,
config=VastbaseVectorConfig(
host="localhost",
port=5434,
user="dify",
password="Difyai123456",
database="dify",
min_connection=1,
max_connection=5,
),
)
def test_vastbase_vector(setup_mock_redis):
VastbaseVectorTest().run_all_tests()
@@ -1,110 +0,0 @@
import uuid
from core.rag.datasource.vdb.qdrant.qdrant_vector import QdrantConfig, QdrantVector
from core.rag.models.document import Document
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class QdrantVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.attributes = ["doc_id", "dataset_id", "document_id", "doc_hash"]
self.vector = QdrantVector(
collection_name=self.collection_name,
group_id=self.dataset_id,
config=QdrantConfig(
endpoint="http://localhost:6333",
api_key="difyai123456",
),
)
# Additional doc IDs for multi-keyword search tests
self.doc_apple_id = ""
self.doc_banana_id = ""
self.doc_both_id = ""
def search_by_vector(self):
super().search_by_vector()
# only test for qdrant, may not work on other vector stores
hits_by_vector: list[Document] = self.vector.search_by_vector(
query_vector=self.example_embedding, score_threshold=1
)
assert len(hits_by_vector) == 0
def _create_document(self, content: str, doc_id: str) -> Document:
"""Create a document with the given content and doc_id."""
return Document(
page_content=content,
metadata={
"doc_id": doc_id,
"doc_hash": doc_id,
"document_id": doc_id,
"dataset_id": self.dataset_id,
},
)
def setup_multi_keyword_documents(self):
"""Create test documents with different keyword combinations for multi-keyword search tests."""
self.doc_apple_id = str(uuid.uuid4())
self.doc_banana_id = str(uuid.uuid4())
self.doc_both_id = str(uuid.uuid4())
documents = [
self._create_document("This document contains apple only", self.doc_apple_id),
self._create_document("This document contains banana only", self.doc_banana_id),
self._create_document("This document contains both apple and banana", self.doc_both_id),
]
embeddings = [self.example_embedding] * len(documents)
self.vector.add_texts(documents=documents, embeddings=embeddings)
def search_by_full_text_multi_keyword(self):
"""Test multi-keyword search returns docs matching ANY keyword (OR logic)."""
# First verify single keyword searches work correctly
hits_apple = self.vector.search_by_full_text(query="apple", top_k=10)
apple_ids = {doc.metadata["doc_id"] for doc in hits_apple}
assert self.doc_apple_id in apple_ids, "Document with 'apple' should be found"
assert self.doc_both_id in apple_ids, "Document with 'apple and banana' should be found"
hits_banana = self.vector.search_by_full_text(query="banana", top_k=10)
banana_ids = {doc.metadata["doc_id"] for doc in hits_banana}
assert self.doc_banana_id in banana_ids, "Document with 'banana' should be found"
assert self.doc_both_id in banana_ids, "Document with 'apple and banana' should be found"
# Test multi-keyword search returns all matching documents
hits = self.vector.search_by_full_text(query="apple banana", top_k=10)
doc_ids = {doc.metadata["doc_id"] for doc in hits}
assert self.doc_apple_id in doc_ids, "Document with 'apple' should be found in multi-keyword search"
assert self.doc_banana_id in doc_ids, "Document with 'banana' should be found in multi-keyword search"
assert self.doc_both_id in doc_ids, "Document with both keywords should be found"
# Expect 3 results: doc_apple (apple only), doc_banana (banana only), doc_both (contains both)
assert len(hits) == 3, f"Expected 3 documents, got {len(hits)}"
# Test keyword order independence
hits_ba = self.vector.search_by_full_text(query="banana apple", top_k=10)
ids_ba = {doc.metadata["doc_id"] for doc in hits_ba}
assert doc_ids == ids_ba, "Keyword order should not affect search results"
# Test no duplicates in results
doc_id_list = [doc.metadata["doc_id"] for doc in hits]
assert len(doc_id_list) == len(set(doc_id_list)), "Search results should not contain duplicates"
def run_all_tests(self):
self.create_vector()
self.search_by_vector()
self.search_by_full_text()
self.text_exists()
self.get_ids_by_metadata_field()
# Multi-keyword search tests
self.setup_multi_keyword_documents()
self.search_by_full_text_multi_keyword()
# Cleanup - delete_vector() removes the entire collection
self.delete_vector()
def test_qdrant_vector(setup_mock_redis):
QdrantVectorTest().run_all_tests()
@@ -1,101 +0,0 @@
import os
import uuid
import tablestore
from _pytest.python_api import approx
from core.rag.datasource.vdb.tablestore.tablestore_vector import (
TableStoreConfig,
TableStoreVector,
)
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
get_example_document,
get_example_text,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class TableStoreVectorTest(AbstractVectorTest):
def __init__(self, normalize_full_text_score: bool = False):
super().__init__()
self.vector = TableStoreVector(
collection_name=self.collection_name,
config=TableStoreConfig(
endpoint=os.getenv("TABLESTORE_ENDPOINT"),
instance_name=os.getenv("TABLESTORE_INSTANCE_NAME"),
access_key_id=os.getenv("TABLESTORE_ACCESS_KEY_ID"),
access_key_secret=os.getenv("TABLESTORE_ACCESS_KEY_SECRET"),
normalize_full_text_bm25_score=normalize_full_text_score,
),
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="doc_id", value=self.example_doc_id)
assert ids is not None
assert len(ids) == 1
assert ids[0] == self.example_doc_id
def create_vector(self):
self.vector.create(
texts=[get_example_document(doc_id=self.example_doc_id)],
embeddings=[self.example_embedding],
)
while True:
search_response = self.vector._tablestore_client.search(
table_name=self.vector._table_name,
index_name=self.vector._index_name,
search_query=tablestore.SearchQuery(query=tablestore.MatchAllQuery(), get_total_count=True, limit=0),
columns_to_get=tablestore.ColumnsToGet(return_type=tablestore.ColumnReturnType.ALL_FROM_INDEX),
)
if search_response.total_count == 1:
break
def search_by_vector(self):
super().search_by_vector()
docs = self.vector.search_by_vector(self.example_embedding, document_ids_filter=[self.example_doc_id])
assert len(docs) == 1
assert docs[0].metadata["doc_id"] == self.example_doc_id
assert docs[0].metadata["score"] > 0
docs = self.vector.search_by_vector(self.example_embedding, document_ids_filter=[str(uuid.uuid4())])
assert len(docs) == 0
def search_by_full_text(self):
super().search_by_full_text()
docs = self.vector.search_by_full_text(get_example_text(), document_ids_filter=[self.example_doc_id])
assert len(docs) == 1
assert docs[0].metadata["doc_id"] == self.example_doc_id
if self.vector._config.normalize_full_text_bm25_score:
assert docs[0].metadata["score"] == approx(0.1214, abs=1e-3)
else:
assert docs[0].metadata.get("score") is None
# return none if normalize_full_text_score=true and score_threshold > 0
docs = self.vector.search_by_full_text(
get_example_text(), document_ids_filter=[self.example_doc_id], score_threshold=0.5
)
if self.vector._config.normalize_full_text_bm25_score:
assert len(docs) == 0
else:
assert len(docs) == 1
assert docs[0].metadata["doc_id"] == self.example_doc_id
assert docs[0].metadata.get("score") is None
docs = self.vector.search_by_full_text(get_example_text(), document_ids_filter=[str(uuid.uuid4())])
assert len(docs) == 0
def run_all_tests(self):
try:
self.vector.delete()
except Exception:
pass
return super().run_all_tests()
def test_tablestore_vector(setup_mock_redis):
TableStoreVectorTest().run_all_tests()
TableStoreVectorTest(normalize_full_text_score=True).run_all_tests()
TableStoreVectorTest(normalize_full_text_score=False).run_all_tests()
@@ -1,42 +0,0 @@
from unittest.mock import MagicMock
from core.rag.datasource.vdb.tencent.tencent_vector import TencentConfig, TencentVector
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = (
"tests.integration_tests.vdb.test_vector_store",
"tests.integration_tests.vdb.__mock.tcvectordb",
)
mock_client = MagicMock()
mock_client.list_databases.return_value = [{"name": "test"}]
class TencentVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = TencentVector(
"dify",
TencentConfig(
url="http://127.0.0.1",
api_key="dify",
timeout=30,
username="dify",
database="dify",
shard=1,
replicas=2,
enable_hybrid_search=True,
),
)
def search_by_vector(self):
hits_by_vector = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 1
def search_by_full_text(self):
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) >= 0
def test_tencent_vector(setup_mock_redis, setup_tcvectordb_mock):
TencentVectorTest().run_all_tests()
@@ -1,95 +0,0 @@
import uuid
from unittest.mock import MagicMock
import pytest
from core.rag.models.document import Document
from extensions import ext_redis
from models.dataset import Dataset
def get_example_text() -> str:
return "test_text"
def get_example_document(doc_id: str) -> Document:
doc = Document(
page_content=get_example_text(),
metadata={
"doc_id": doc_id,
"doc_hash": doc_id,
"document_id": doc_id,
"dataset_id": doc_id,
},
)
return doc
@pytest.fixture
def setup_mock_redis():
# get
ext_redis.redis_client.get = MagicMock(return_value=None)
# set
ext_redis.redis_client.set = MagicMock(return_value=None)
# lock
mock_redis_lock = MagicMock()
mock_redis_lock.__enter__ = MagicMock()
mock_redis_lock.__exit__ = MagicMock()
ext_redis.redis_client.lock = mock_redis_lock
class AbstractVectorTest:
def __init__(self):
self.vector = None
self.dataset_id = str(uuid.uuid4())
self.collection_name = Dataset.gen_collection_name_by_id(self.dataset_id) + "_test"
self.example_doc_id = str(uuid.uuid4())
self.example_embedding = [1.001 * i for i in range(128)]
def create_vector(self):
self.vector.create(
texts=[get_example_document(doc_id=self.example_doc_id)],
embeddings=[self.example_embedding],
)
def search_by_vector(self):
hits_by_vector: list[Document] = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 1
assert hits_by_vector[0].metadata["doc_id"] == self.example_doc_id
def search_by_full_text(self):
hits_by_full_text: list[Document] = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 1
assert hits_by_full_text[0].metadata["doc_id"] == self.example_doc_id
def delete_vector(self):
self.vector.delete()
def delete_by_ids(self, ids: list[str]):
self.vector.delete_by_ids(ids=ids)
def add_texts(self) -> list[str]:
batch_size = 100
documents = [get_example_document(doc_id=str(uuid.uuid4())) for _ in range(batch_size)]
embeddings = [self.example_embedding] * batch_size
self.vector.add_texts(documents=documents, embeddings=embeddings)
return [doc.metadata["doc_id"] for doc in documents]
def text_exists(self):
assert self.vector.text_exists(self.example_doc_id)
def get_ids_by_metadata_field(self):
with pytest.raises(NotImplementedError):
self.vector.get_ids_by_metadata_field(key="key", value="value")
def run_all_tests(self):
self.create_vector()
self.search_by_vector()
self.search_by_full_text()
self.text_exists()
self.get_ids_by_metadata_field()
added_doc_ids = self.add_texts()
self.delete_by_ids(added_doc_ids)
self.delete_vector()
@@ -1,59 +0,0 @@
import time
import pymysql
def check_tiflash_ready() -> bool:
try:
connection = pymysql.connect(
host="localhost",
port=4000,
user="root",
password="",
)
with connection.cursor() as cursor:
# Doc reference:
# https://docs.pingcap.com/zh/tidb/stable/information-schema-cluster-hardware
select_tiflash_query = """
SELECT * FROM information_schema.cluster_hardware
WHERE TYPE='tiflash'
LIMIT 1;
"""
cursor.execute(select_tiflash_query)
result = cursor.fetchall()
return result is not None and len(result) > 0
except Exception as e:
print(f"TiFlash is not ready. Exception: {e}")
return False
finally:
if connection:
connection.close()
def main():
max_attempts = 30
retry_interval_seconds = 2
is_tiflash_ready = False
for attempt in range(max_attempts):
try:
is_tiflash_ready = check_tiflash_ready()
except Exception as e:
print(f"TiFlash is not ready. Exception: {e}")
is_tiflash_ready = False
if is_tiflash_ready:
break
else:
print(f"Attempt {attempt + 1} failed, retry in {retry_interval_seconds} seconds...")
time.sleep(retry_interval_seconds)
if is_tiflash_ready:
print("TiFlash is ready in TiDB.")
else:
print(f"TiFlash is not ready in TiDB after {max_attempts} attempting checks.")
exit(1)
if __name__ == "__main__":
main()
@@ -1,42 +0,0 @@
import pytest
from core.rag.datasource.vdb.tidb_vector.tidb_vector import TiDBVector, TiDBVectorConfig
from models.dataset import Document
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
@pytest.fixture
def tidb_vector():
return TiDBVector(
collection_name="test_collection",
config=TiDBVectorConfig(
host="localhost",
port=4000,
user="root",
password="",
database="test",
program_name="langgenius/dify",
),
)
class TiDBVectorTest(AbstractVectorTest):
def __init__(self, vector):
super().__init__()
self.vector = vector
def search_by_full_text(self):
hits_by_full_text: list[Document] = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="doc_id", value=self.example_doc_id)
assert len(ids) == 1
def test_tidb_vector(setup_mock_redis, tidb_vector):
# TiDBVectorTest(vector=tidb_vector).run_all_tests()
# something wrong with tidb,ignore tidb test
return
@@ -1,29 +0,0 @@
from core.rag.datasource.vdb.upstash.upstash_vector import UpstashVector, UpstashVectorConfig
from core.rag.models.document import Document
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = ("tests.integration_tests.vdb.__mock.upstashvectordb",)
class UpstashVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = UpstashVector(
collection_name="test_collection",
config=UpstashVectorConfig(
url="your-server-url",
token="your-access-token",
),
)
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value=self.example_doc_id)
assert len(ids) != 0
def search_by_full_text(self):
hits_by_full_text: list[Document] = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def test_upstash_vector(setup_upstashvector_mock):
UpstashVectorTest().run_all_tests()
@@ -1,41 +0,0 @@
from core.rag.datasource.vdb.vikingdb.vikingdb_vector import VikingDBConfig, VikingDBVector
from tests.integration_tests.vdb.test_vector_store import AbstractVectorTest, get_example_text
pytest_plugins = (
"tests.integration_tests.vdb.test_vector_store",
"tests.integration_tests.vdb.__mock.vikingdb",
)
class VikingDBVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.vector = VikingDBVector(
"test_collection",
"test_group",
config=VikingDBConfig(
access_key="test_access_key",
host="test_host",
region="test_region",
scheme="test_scheme",
secret_key="test_secret_key",
connection_timeout=30,
socket_timeout=30,
),
)
def search_by_vector(self):
hits_by_vector = self.vector.search_by_vector(query_vector=self.example_embedding)
assert len(hits_by_vector) == 1
def search_by_full_text(self):
hits_by_full_text = self.vector.search_by_full_text(query=get_example_text())
assert len(hits_by_full_text) == 0
def get_ids_by_metadata_field(self):
ids = self.vector.get_ids_by_metadata_field(key="document_id", value="test_document_id")
assert len(ids) > 0
def test_vikingdb_vector(setup_mock_redis, setup_vikingdb_mock):
VikingDBVectorTest().run_all_tests()
@@ -1,24 +0,0 @@
from core.rag.datasource.vdb.weaviate.weaviate_vector import WeaviateConfig, WeaviateVector
from tests.integration_tests.vdb.test_vector_store import (
AbstractVectorTest,
)
pytest_plugins = ("tests.integration_tests.vdb.test_vector_store",)
class WeaviateVectorTest(AbstractVectorTest):
def __init__(self):
super().__init__()
self.attributes = ["doc_id", "dataset_id", "document_id", "doc_hash"]
self.vector = WeaviateVector(
collection_name=self.collection_name,
config=WeaviateConfig(
endpoint="http://localhost:8080",
api_key="WVF5YThaHlkYwhGUSmCRgsX3tD5ngdN8pkih",
),
attributes=self.attributes,
)
def test_weaviate_vector(setup_mock_redis):
WeaviateVectorTest().run_all_tests()
@@ -1,74 +0,0 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector as alibaba_module
from core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector import AlibabaCloudMySQLVectorFactory
def test_validate_distance_function_accepts_supported_values():
factory = AlibabaCloudMySQLVectorFactory()
assert factory._validate_distance_function("cosine") == "cosine"
assert factory._validate_distance_function("euclidean") == "euclidean"
def test_validate_distance_function_rejects_unsupported_values():
factory = AlibabaCloudMySQLVectorFactory()
with pytest.raises(ValueError, match="Invalid distance function"):
factory._validate_distance_function("dot_product")
def test_factory_init_vector_uses_existing_index_struct_class_prefix(monkeypatch):
factory = AlibabaCloudMySQLVectorFactory()
dataset = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "existing_collection"}},
index_struct=None,
)
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_HOST", "host")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_PORT", 3306)
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_USER", "user")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_PASSWORD", "password")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_DATABASE", "db")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_MAX_CONNECTION", 5)
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_CHARSET", "utf8mb4")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_DISTANCE_FUNCTION", "cosine")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_HNSW_M", 6)
with patch.object(alibaba_module, "AlibabaCloudMySQLVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
assert vector_cls.call_args.kwargs["collection_name"] == "existing_collection"
def test_factory_init_vector_generates_collection_name_when_index_struct_is_missing(monkeypatch):
factory = AlibabaCloudMySQLVectorFactory()
dataset = SimpleNamespace(
id="dataset-2",
index_struct_dict=None,
index_struct=None,
)
monkeypatch.setattr(alibaba_module.Dataset, "gen_collection_name_by_id", lambda dataset_id: f"COL_{dataset_id}")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_HOST", "host")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_PORT", 3306)
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_USER", "user")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_PASSWORD", "password")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_DATABASE", "db")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_MAX_CONNECTION", 5)
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_CHARSET", "utf8mb4")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_DISTANCE_FUNCTION", "euclidean")
monkeypatch.setattr(alibaba_module.dify_config, "ALIBABACLOUD_MYSQL_HNSW_M", 12)
with patch.object(alibaba_module, "AlibabaCloudMySQLVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
vector_cls.assert_called_once()
assert vector_cls.call_args.kwargs["collection_name"] == "COL_dataset-2"
assert dataset.index_struct is not None
@@ -1,733 +0,0 @@
import json
import unittest
from unittest.mock import MagicMock, patch
import pytest
from core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector import (
AlibabaCloudMySQLVector,
AlibabaCloudMySQLVectorConfig,
)
from core.rag.models.document import Document
try:
from mysql.connector import Error as MySQLError
except ImportError:
# Fallback for testing environments where mysql-connector-python might not be installed
class MySQLError(Exception):
def __init__(self, errno, msg):
self.errno = errno
self.msg = msg
super().__init__(msg)
class TestAlibabaCloudMySQLVector(unittest.TestCase):
def setUp(self):
self.config = AlibabaCloudMySQLVectorConfig(
host="localhost",
port=3306,
user="test_user",
password="test_password",
database="test_db",
max_connection=5,
charset="utf8mb4",
)
self.collection_name = "test_collection"
# Sample documents for testing
self.sample_documents = [
Document(
page_content="This is a test document about AI.",
metadata={"doc_id": "doc1", "document_id": "dataset1", "source": "test"},
),
Document(
page_content="Another document about machine learning.",
metadata={"doc_id": "doc2", "document_id": "dataset1", "source": "test"},
),
]
# Sample embeddings
self.sample_embeddings = [[0.1, 0.2, 0.3, 0.4], [0.5, 0.6, 0.7, 0.8]]
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_init(self, mock_pool_class):
"""Test AlibabaCloudMySQLVector initialization."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor for vector support check
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [
{"VERSION()": "8.0.36"}, # Version check
{"vector_support": True}, # Vector support check
]
alibabacloud_mysql_vector = AlibabaCloudMySQLVector(self.collection_name, self.config)
assert alibabacloud_mysql_vector.collection_name == self.collection_name
assert alibabacloud_mysql_vector.table_name == self.collection_name.lower()
assert alibabacloud_mysql_vector.get_type() == "alibabacloud_mysql"
assert alibabacloud_mysql_vector.distance_function == "cosine"
assert alibabacloud_mysql_vector.pool is not None
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
@patch("core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.redis_client")
def test_create_collection(self, mock_redis, mock_pool_class):
"""Test collection creation."""
# Mock Redis operations
mock_redis.lock.return_value.__enter__ = MagicMock()
mock_redis.lock.return_value.__exit__ = MagicMock()
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [
{"VERSION()": "8.0.36"}, # Version check
{"vector_support": True}, # Vector support check
]
alibabacloud_mysql_vector = AlibabaCloudMySQLVector(self.collection_name, self.config)
alibabacloud_mysql_vector._create_collection(768)
# Verify SQL execution calls - should include table creation and index creation
assert mock_cursor.execute.called
assert mock_cursor.execute.call_count >= 3 # CREATE TABLE + 2 indexes
mock_redis.set.assert_called_once()
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_vector_support_check_success(self, mock_pool_class):
"""Test successful vector support check."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
# Should not raise an exception
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
assert vector_store is not None
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_vector_support_check_failure(self, mock_pool_class):
"""Test vector support check failure."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.35"}, {"vector_support": False}]
with pytest.raises(ValueError) as context:
AlibabaCloudMySQLVector(self.collection_name, self.config)
assert "RDS MySQL Vector functions are not available" in str(context.value)
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_vector_support_check_function_error(self, mock_pool_class):
"""Test vector support check with function not found error."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.return_value = {"VERSION()": "8.0.36"}
mock_cursor.execute.side_effect = [None, MySQLError(errno=1305, msg="FUNCTION VEC_FromText does not exist")]
with pytest.raises(ValueError) as context:
AlibabaCloudMySQLVector(self.collection_name, self.config)
assert "RDS MySQL Vector functions are not available" in str(context.value)
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
@patch("core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.redis_client")
def test_create_documents(self, mock_redis, mock_pool_class):
"""Test creating documents with embeddings."""
# Setup mocks
self._setup_mocks(mock_redis, mock_pool_class)
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
result = vector_store.create(self.sample_documents, self.sample_embeddings)
assert len(result) == 2
assert "doc1" in result
assert "doc2" in result
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_add_texts(self, mock_pool_class):
"""Test adding texts to the vector store."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
result = vector_store.add_texts(self.sample_documents, self.sample_embeddings)
assert len(result) == 2
mock_cursor.executemany.assert_called_once()
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_text_exists(self, mock_pool_class):
"""Test checking if text exists."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [
{"VERSION()": "8.0.36"},
{"vector_support": True},
{"id": "doc1"}, # Text exists
]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
exists = vector_store.text_exists("doc1")
assert exists
# Check that the correct SQL was executed (last call after init)
execute_calls = mock_cursor.execute.call_args_list
last_call = execute_calls[-1]
assert "SELECT id FROM" in last_call[0][0]
assert last_call[0][1] == ("doc1",)
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_text_not_exists(self, mock_pool_class):
"""Test checking if text does not exist."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [
{"VERSION()": "8.0.36"},
{"vector_support": True},
None, # Text does not exist
]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
exists = vector_store.text_exists("nonexistent")
assert not exists
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_get_by_ids(self, mock_pool_class):
"""Test getting documents by IDs."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter(
[
{"meta": json.dumps({"doc_id": "doc1", "source": "test"}), "text": "Test document 1"},
{"meta": json.dumps({"doc_id": "doc2", "source": "test"}), "text": "Test document 2"},
]
)
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
docs = vector_store.get_by_ids(["doc1", "doc2"])
assert len(docs) == 2
assert docs[0].page_content == "Test document 1"
assert docs[1].page_content == "Test document 2"
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_get_by_ids_empty_list(self, mock_pool_class):
"""Test getting documents with empty ID list."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
docs = vector_store.get_by_ids([])
assert len(docs) == 0
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_delete_by_ids(self, mock_pool_class):
"""Test deleting documents by IDs."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
vector_store.delete_by_ids(["doc1", "doc2"])
# Check that delete SQL was executed
execute_calls = mock_cursor.execute.call_args_list
delete_calls = [call for call in execute_calls if "DELETE" in str(call)]
assert len(delete_calls) == 1
delete_call = delete_calls[0]
assert "DELETE FROM" in delete_call[0][0]
assert delete_call[0][1] == ["doc1", "doc2"]
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_delete_by_ids_empty_list(self, mock_pool_class):
"""Test deleting with empty ID list."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
vector_store.delete_by_ids([]) # Should not raise an exception
# Verify no delete SQL was executed
execute_calls = mock_cursor.execute.call_args_list
delete_calls = [call for call in execute_calls if "DELETE" in str(call)]
assert len(delete_calls) == 0
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_delete_by_ids_table_not_exists(self, mock_pool_class):
"""Test deleting when table doesn't exist."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
# Simulate table doesn't exist error on delete
def execute_side_effect(*args, **kwargs):
if "DELETE" in args[0]:
raise MySQLError(errno=1146, msg="Table doesn't exist")
mock_cursor.execute.side_effect = execute_side_effect
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
# Should not raise an exception
vector_store.delete_by_ids(["doc1"])
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_delete_by_metadata_field(self, mock_pool_class):
"""Test deleting documents by metadata field."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
vector_store.delete_by_metadata_field("document_id", "dataset1")
# Check that the correct SQL was executed
execute_calls = mock_cursor.execute.call_args_list
delete_calls = [call for call in execute_calls if "DELETE" in str(call)]
assert len(delete_calls) == 1
delete_call = delete_calls[0]
assert "JSON_UNQUOTE(JSON_EXTRACT(meta" in delete_call[0][0]
assert delete_call[0][1] == ("$.document_id", "dataset1")
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_vector_cosine(self, mock_pool_class):
"""Test vector search with cosine distance."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter(
[{"meta": json.dumps({"doc_id": "doc1", "source": "test"}), "text": "Test document 1", "distance": 0.1}]
)
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
query_vector = [0.1, 0.2, 0.3, 0.4]
docs = vector_store.search_by_vector(query_vector, top_k=5)
assert len(docs) == 1
assert docs[0].page_content == "Test document 1"
assert abs(docs[0].metadata["score"] - 0.9) < 0.1 # 1 - 0.1 = 0.9
assert docs[0].metadata["distance"] == 0.1
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_vector_euclidean(self, mock_pool_class):
"""Test vector search with euclidean distance."""
config = AlibabaCloudMySQLVectorConfig(
host="localhost",
port=3306,
user="test_user",
password="test_password",
database="test_db",
max_connection=5,
distance_function="euclidean",
)
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter(
[{"meta": json.dumps({"doc_id": "doc1", "source": "test"}), "text": "Test document 1", "distance": 2.0}]
)
vector_store = AlibabaCloudMySQLVector(self.collection_name, config)
query_vector = [0.1, 0.2, 0.3, 0.4]
docs = vector_store.search_by_vector(query_vector, top_k=5)
assert len(docs) == 1
assert abs(docs[0].metadata["score"] - 1.0 / 3.0) < 0.01 # 1/(1+2) = 1/3
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_vector_with_filter(self, mock_pool_class):
"""Test vector search with document ID filter."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter([])
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
query_vector = [0.1, 0.2, 0.3, 0.4]
docs = vector_store.search_by_vector(query_vector, top_k=5, document_ids_filter=["dataset1"])
# Verify the SQL contains the WHERE clause for filtering
execute_calls = mock_cursor.execute.call_args_list
search_calls = [call for call in execute_calls if "VEC_DISTANCE" in str(call)]
assert len(search_calls) > 0
search_call = search_calls[0]
assert "WHERE JSON_UNQUOTE" in search_call[0][0]
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_vector_with_score_threshold(self, mock_pool_class):
"""Test vector search with score threshold."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter(
[
{
"meta": json.dumps({"doc_id": "doc1", "source": "test"}),
"text": "High similarity document",
"distance": 0.1, # High similarity (score = 0.9)
},
{
"meta": json.dumps({"doc_id": "doc2", "source": "test"}),
"text": "Low similarity document",
"distance": 0.8, # Low similarity (score = 0.2)
},
]
)
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
query_vector = [0.1, 0.2, 0.3, 0.4]
docs = vector_store.search_by_vector(query_vector, top_k=5, score_threshold=0.5)
# Only the high similarity document should be returned
assert len(docs) == 1
assert docs[0].page_content == "High similarity document"
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_vector_invalid_top_k(self, mock_pool_class):
"""Test vector search with invalid top_k."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
query_vector = [0.1, 0.2, 0.3, 0.4]
with pytest.raises(ValueError):
vector_store.search_by_vector(query_vector, top_k=0)
with pytest.raises(ValueError):
vector_store.search_by_vector(query_vector, top_k="invalid")
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_full_text(self, mock_pool_class):
"""Test full-text search."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter(
[
{
"meta": {"doc_id": "doc1", "source": "test"},
"text": "This document contains machine learning content",
"score": 1.5,
}
]
)
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
docs = vector_store.search_by_full_text("machine learning", top_k=5)
assert len(docs) == 1
assert docs[0].page_content == "This document contains machine learning content"
assert docs[0].metadata["score"] == 1.5
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_full_text_with_filter(self, mock_pool_class):
"""Test full-text search with document ID filter."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
mock_cursor.__iter__ = lambda self: iter([])
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
docs = vector_store.search_by_full_text("machine learning", top_k=5, document_ids_filter=["dataset1"])
# Verify the SQL contains the AND clause for filtering
execute_calls = mock_cursor.execute.call_args_list
search_calls = [call for call in execute_calls if "MATCH" in str(call)]
assert len(search_calls) > 0
search_call = search_calls[0]
assert "AND JSON_UNQUOTE" in search_call[0][0]
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_search_by_full_text_invalid_top_k(self, mock_pool_class):
"""Test full-text search with invalid top_k."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
with pytest.raises(ValueError):
vector_store.search_by_full_text("test", top_k=0)
with pytest.raises(ValueError):
vector_store.search_by_full_text("test", top_k="invalid")
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_delete_collection(self, mock_pool_class):
"""Test deleting the entire collection."""
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
vector_store = AlibabaCloudMySQLVector(self.collection_name, self.config)
vector_store.delete()
# Check that DROP TABLE SQL was executed
execute_calls = mock_cursor.execute.call_args_list
drop_calls = [call for call in execute_calls if "DROP TABLE" in str(call)]
assert len(drop_calls) == 1
drop_call = drop_calls[0]
assert f"DROP TABLE IF EXISTS {self.collection_name.lower()}" in drop_call[0][0]
@patch(
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector.mysql.connector.pooling.MySQLConnectionPool"
)
def test_unsupported_distance_function(self, mock_pool_class):
"""Test that Pydantic validation rejects unsupported distance functions."""
# Test that creating config with unsupported distance function raises ValidationError
with pytest.raises(ValueError) as context:
AlibabaCloudMySQLVectorConfig(
host="localhost",
port=3306,
user="test_user",
password="test_password",
database="test_db",
max_connection=5,
distance_function="manhattan", # Unsupported - not in Literal["cosine", "euclidean"]
)
# The error should be related to validation
assert "Input should be 'cosine' or 'euclidean'" in str(context.value) or "manhattan" in str(context.value)
def _setup_mocks(self, mock_redis, mock_pool_class):
"""Helper method to setup common mocks."""
# Mock Redis operations
mock_redis.lock.return_value.__enter__ = MagicMock()
mock_redis.lock.return_value.__exit__ = MagicMock()
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.get_connection.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.side_effect = [{"VERSION()": "8.0.36"}, {"vector_support": True}]
@pytest.mark.parametrize(
"invalid_config_override",
[
{"host": ""}, # Test empty host
{"port": 0}, # Test invalid port
{"max_connection": 0}, # Test invalid max_connection
],
)
def test_config_validation_parametrized(invalid_config_override):
"""Test configuration validation for various invalid inputs using parametrize."""
config = {
"host": "localhost",
"port": 3306,
"user": "test",
"password": "test",
"database": "test",
"max_connection": 5,
}
config.update(invalid_config_override)
with pytest.raises(ValueError):
AlibabaCloudMySQLVectorConfig(**config)
if __name__ == "__main__":
unittest.main()
@@ -1,133 +0,0 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import core.rag.datasource.vdb.analyticdb.analyticdb_vector as analyticdb_module
from core.rag.datasource.vdb.analyticdb.analyticdb_vector import AnalyticdbVector, AnalyticdbVectorFactory
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_openapi import AnalyticdbVectorOpenAPIConfig
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_sql import AnalyticdbVectorBySqlConfig
from core.rag.models.document import Document
def test_init_prefers_openapi_when_api_config_is_provided():
api_config = AnalyticdbVectorOpenAPIConfig(
access_key_id="ak",
access_key_secret="sk",
region_id="cn-hangzhou",
instance_id="instance-1",
account="account",
account_password="password",
namespace="dify",
namespace_password="ns-password",
)
with patch.object(analyticdb_module, "AnalyticdbVectorOpenAPI", return_value="openapi_runner") as openapi_cls:
vector = AnalyticdbVector("COLLECTION", api_config=api_config, sql_config=None)
assert vector.analyticdb_vector == "openapi_runner"
openapi_cls.assert_called_once_with("COLLECTION", api_config)
def test_init_uses_sql_implementation_when_api_config_is_missing():
sql_config = AnalyticdbVectorBySqlConfig(
host="localhost",
port=5432,
account="account",
account_password="password",
min_connection=1,
max_connection=2,
namespace="dify",
)
with patch.object(analyticdb_module, "AnalyticdbVectorBySql", return_value="sql_runner") as sql_cls:
vector = AnalyticdbVector("COLLECTION", api_config=None, sql_config=sql_config)
assert vector.analyticdb_vector == "sql_runner"
sql_cls.assert_called_once_with("COLLECTION", sql_config)
def test_init_raises_when_both_configs_are_missing():
with pytest.raises(ValueError, match="Either api_config or sql_config must be provided"):
AnalyticdbVector("COLLECTION", api_config=None, sql_config=None)
def test_vector_methods_delegate_to_underlying_implementation():
runner = MagicMock()
runner.search_by_vector.return_value = [Document(page_content="v", metadata={"doc_id": "1"})]
runner.search_by_full_text.return_value = [Document(page_content="t", metadata={"doc_id": "2"})]
runner.text_exists.return_value = True
vector = AnalyticdbVector.__new__(AnalyticdbVector)
vector.analyticdb_vector = runner
texts = [Document(page_content="hello", metadata={"doc_id": "d1"})]
vector.create(texts=texts, embeddings=[[0.1, 0.2]])
vector.add_texts(documents=texts, embeddings=[[0.1, 0.2]])
assert vector.text_exists("d1") is True
vector.delete_by_ids(["d1"])
vector.delete_by_metadata_field("document_id", "doc-1")
assert vector.search_by_vector([0.1, 0.2], top_k=2) == runner.search_by_vector.return_value
assert vector.search_by_full_text("hello", top_k=2) == runner.search_by_full_text.return_value
vector.delete()
runner.create_collection_if_not_exists.assert_called_once_with(2)
runner.add_texts.assert_any_call(texts, [[0.1, 0.2]])
runner.delete_by_ids.assert_called_once_with(["d1"])
runner.delete_by_metadata_field.assert_called_once_with("document_id", "doc-1")
runner.delete.assert_called_once()
def test_get_type_is_analyticdb():
vector = AnalyticdbVector.__new__(AnalyticdbVector)
assert vector.get_type() == "analyticdb"
def test_factory_builds_openapi_config_when_host_is_missing(monkeypatch):
factory = AnalyticdbVectorFactory()
dataset = SimpleNamespace(id="dataset-1", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(analyticdb_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_HOST", None)
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_KEY_ID", "ak")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_KEY_SECRET", "sk")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_REGION_ID", "cn-hz")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_INSTANCE_ID", "instance")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_ACCOUNT", "account")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_PASSWORD", "password")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_NAMESPACE", "dify")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_NAMESPACE_PASSWORD", "ns-password")
with patch.object(analyticdb_module, "AnalyticdbVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
args = vector_cls.call_args.args
assert args[0] == "auto_collection"
assert isinstance(args[1], AnalyticdbVectorOpenAPIConfig)
assert args[2] is None
assert dataset.index_struct is not None
def test_factory_builds_sql_config_when_host_is_present(monkeypatch):
factory = AnalyticdbVectorFactory()
dataset = SimpleNamespace(
id="dataset-2", index_struct_dict={"vector_store": {"class_prefix": "EXISTING"}}, index_struct=None
)
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_HOST", "127.0.0.1")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_PORT", 5432)
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_ACCOUNT", "account")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_PASSWORD", "password")
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_MIN_CONNECTION", 1)
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_MAX_CONNECTION", 3)
monkeypatch.setattr(analyticdb_module.dify_config, "ANALYTICDB_NAMESPACE", "dify")
with patch.object(analyticdb_module, "AnalyticdbVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
args = vector_cls.call_args.args
assert args[0] == "existing"
assert args[1] is None
assert isinstance(args[2], AnalyticdbVectorBySqlConfig)
@@ -1,384 +0,0 @@
import json
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
import core.rag.datasource.vdb.analyticdb.analyticdb_vector_openapi as openapi_module
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_openapi import (
AnalyticdbVectorOpenAPI,
AnalyticdbVectorOpenAPIConfig,
)
from core.rag.models.document import Document
def _request_class(name: str):
class _Request:
def __init__(self, **kwargs):
for key, value in kwargs.items():
setattr(self, key, value)
_Request.__name__ = name
return _Request
def _install_openapi_stubs(monkeypatch):
gpdb_package = types.ModuleType("alibabacloud_gpdb20160503")
gpdb_package.__path__ = []
gpdb_models = types.ModuleType("alibabacloud_gpdb20160503.models")
for class_name in [
"InitVectorDatabaseRequest",
"DescribeNamespaceRequest",
"CreateNamespaceRequest",
"DescribeCollectionRequest",
"CreateCollectionRequest",
"UpsertCollectionDataRequestRows",
"UpsertCollectionDataRequest",
"QueryCollectionDataRequest",
"DeleteCollectionDataRequest",
"DeleteCollectionRequest",
]:
setattr(gpdb_models, class_name, _request_class(class_name))
class _Client:
def __init__(self, config):
self.config = config
gpdb_client = types.ModuleType("alibabacloud_gpdb20160503.client")
gpdb_client.Client = _Client
gpdb_package.models = gpdb_models
tea_openapi = types.ModuleType("alibabacloud_tea_openapi")
tea_openapi.__path__ = []
tea_openapi_models = types.ModuleType("alibabacloud_tea_openapi.models")
class OpenApiConfig:
def __init__(self, **kwargs):
for key, value in kwargs.items():
setattr(self, key, value)
tea_openapi_models.Config = OpenApiConfig
tea_openapi.models = tea_openapi_models
tea_package = types.ModuleType("Tea")
tea_package.__path__ = []
tea_exceptions = types.ModuleType("Tea.exceptions")
class TeaError(Exception):
def __init__(self, status_code=None, **kwargs):
super().__init__("TeaException")
status_code = kwargs.get("statusCode", status_code)
self.statusCode = status_code
self.status_code = status_code
tea_exceptions.TeaException = TeaError
tea_package.exceptions = tea_exceptions
monkeypatch.setitem(sys.modules, "alibabacloud_gpdb20160503", gpdb_package)
monkeypatch.setitem(sys.modules, "alibabacloud_gpdb20160503.models", gpdb_models)
monkeypatch.setitem(sys.modules, "alibabacloud_gpdb20160503.client", gpdb_client)
monkeypatch.setitem(sys.modules, "alibabacloud_tea_openapi", tea_openapi)
monkeypatch.setitem(sys.modules, "alibabacloud_tea_openapi.models", tea_openapi_models)
monkeypatch.setitem(sys.modules, "Tea", tea_package)
monkeypatch.setitem(sys.modules, "Tea.exceptions", tea_exceptions)
return SimpleNamespace(models=gpdb_models, TeaException=TeaError, OpenApiConfig=OpenApiConfig)
def _config() -> AnalyticdbVectorOpenAPIConfig:
return AnalyticdbVectorOpenAPIConfig(
access_key_id="ak",
access_key_secret="sk",
region_id="cn-hangzhou",
instance_id="instance-1",
account="account",
account_password="password",
namespace="dify",
namespace_password="ns-password",
)
@pytest.mark.parametrize(
("field", "value", "error_message"),
[
("access_key_id", "", "ANALYTICDB_KEY_ID"),
("access_key_secret", "", "ANALYTICDB_KEY_SECRET"),
("region_id", "", "ANALYTICDB_REGION_ID"),
("instance_id", "", "ANALYTICDB_INSTANCE_ID"),
("account", "", "ANALYTICDB_ACCOUNT"),
("account_password", "", "ANALYTICDB_PASSWORD"),
("namespace_password", "", "ANALYTICDB_NAMESPACE_PASSWORD"),
],
)
def test_openapi_config_validation(field, value, error_message):
values = _config().model_dump()
values[field] = value
with pytest.raises(ValueError, match=error_message):
AnalyticdbVectorOpenAPIConfig.model_validate(values)
def test_openapi_config_to_client_params():
config = _config()
params = config.to_analyticdb_client_params()
assert params["access_key_id"] == "ak"
assert params["access_key_secret"] == "sk"
assert params["region_id"] == "cn-hangzhou"
assert params["read_timeout"] == 60000
def test_init_creates_openapi_client_and_runs_initialize(monkeypatch):
stubs = _install_openapi_stubs(monkeypatch)
initialize_mock = MagicMock()
monkeypatch.setattr(openapi_module.AnalyticdbVectorOpenAPI, "_initialize", initialize_mock)
vector = AnalyticdbVectorOpenAPI("COLLECTION_1", _config())
assert vector._collection_name == "collection_1"
assert isinstance(vector._client_config, stubs.OpenApiConfig)
assert vector._client_config.user_agent == "dify"
assert vector._client_config.access_key_id == "ak"
assert vector._client.config is vector._client_config
initialize_mock.assert_called_once_with()
def test_initialize_skips_when_cached(monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(openapi_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(openapi_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(openapi_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._initialize_vector_database = MagicMock()
vector._create_namespace_if_not_exists = MagicMock()
vector._initialize()
vector._initialize_vector_database.assert_not_called()
vector._create_namespace_if_not_exists.assert_not_called()
def test_initialize_runs_when_cache_is_missing(monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(openapi_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(openapi_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(openapi_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._initialize_vector_database = MagicMock()
vector._create_namespace_if_not_exists = MagicMock()
vector._initialize()
vector._initialize_vector_database.assert_called_once()
vector._create_namespace_if_not_exists.assert_called_once()
openapi_module.redis_client.set.assert_called_once()
def test_initialize_vector_database_calls_openapi_client(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._client = MagicMock()
vector._initialize_vector_database()
request = vector._client.init_vector_database.call_args.args[0]
assert request.dbinstance_id == "instance-1"
assert request.region_id == "cn-hangzhou"
assert request.manager_account == "account"
assert request.manager_account_password == "password"
def test_create_namespace_creates_when_namespace_not_found(monkeypatch):
stubs = _install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._client = MagicMock()
vector._client.describe_namespace.side_effect = stubs.TeaException(statusCode=404)
vector._create_namespace_if_not_exists()
vector._client.create_namespace.assert_called_once()
def test_create_namespace_raises_on_unexpected_api_error(monkeypatch):
stubs = _install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._client = MagicMock()
vector._client.describe_namespace.side_effect = stubs.TeaException(statusCode=500)
with pytest.raises(ValueError, match="failed to create namespace"):
vector._create_namespace_if_not_exists()
def test_create_namespace_noop_when_namespace_exists(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector.config = _config()
vector._client = MagicMock()
vector._create_namespace_if_not_exists()
vector._client.describe_namespace.assert_called_once()
vector._client.create_namespace.assert_not_called()
def test_create_collection_if_not_exists_creates_when_missing(monkeypatch):
stubs = _install_openapi_stubs(monkeypatch)
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(openapi_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(openapi_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(openapi_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector._client.describe_collection.side_effect = stubs.TeaException(statusCode=404)
vector.create_collection_if_not_exists(embedding_dimension=1024)
vector._client.create_collection.assert_called_once()
openapi_module.redis_client.set.assert_called_once()
def test_create_collection_if_not_exists_skips_when_cached(monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(openapi_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(openapi_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(openapi_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector.create_collection_if_not_exists(embedding_dimension=1024)
vector._client.describe_collection.assert_not_called()
vector._client.create_collection.assert_not_called()
def test_create_collection_if_not_exists_raises_on_non_404_errors(monkeypatch):
stubs = _install_openapi_stubs(monkeypatch)
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(openapi_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(openapi_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(openapi_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector._client.describe_collection.side_effect = stubs.TeaException(statusCode=500)
with pytest.raises(ValueError, match="failed to create collection collection_1"):
vector.create_collection_if_not_exists(embedding_dimension=512)
def test_openapi_add_delete_and_search_methods(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
documents = [
Document(page_content="doc 1", metadata={"doc_id": "d1", "document_id": "doc-1"}),
SimpleNamespace(page_content="doc 2", metadata=None),
]
embeddings = [[0.1, 0.2], [0.2, 0.3]]
vector.add_texts(documents, embeddings)
upsert_request = vector._client.upsert_collection_data.call_args.args[0]
assert upsert_request.collection == "collection_1"
assert len(upsert_request.rows) == 1
vector._client.query_collection_data.return_value = SimpleNamespace(
body=SimpleNamespace(matches=SimpleNamespace(match=[SimpleNamespace()]))
)
assert vector.text_exists("d1") is True
vector.delete_by_ids(["d1", "d2"])
request = vector._client.delete_collection_data.call_args.args[0]
assert request.collection_data_filter == "ref_doc_id IN ('d1','d2')"
vector.delete_by_metadata_field("document_id", "doc-1")
request = vector._client.delete_collection_data.call_args.args[0]
assert request.collection_data_filter == "metadata_ ->> 'document_id' = 'doc-1'"
match_high = SimpleNamespace(
score=0.9,
metadata={"metadata_": json.dumps({"document_id": "doc-1"}), "page_content": "high"},
values=SimpleNamespace(value=[1.0, 2.0]),
)
match_low = SimpleNamespace(
score=0.1,
metadata={"metadata_": json.dumps({"document_id": "doc-2"}), "page_content": "low"},
values=SimpleNamespace(value=[3.0, 4.0]),
)
vector._client.query_collection_data.return_value = SimpleNamespace(
body=SimpleNamespace(matches=SimpleNamespace(match=[match_low, match_high]))
)
docs_by_vector = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["doc-1"])
assert len(docs_by_vector) == 1
assert docs_by_vector[0].page_content == "high"
assert docs_by_vector[0].metadata["score"] == 0.9
docs_by_text = vector.search_by_full_text("hello", top_k=2, score_threshold=0.2)
assert len(docs_by_text) == 1
assert docs_by_text[0].page_content == "high"
def test_text_exists_returns_false_when_matches_empty(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector._client.query_collection_data.return_value = SimpleNamespace(
body=SimpleNamespace(matches=SimpleNamespace(match=[]))
)
assert vector.text_exists("missing-id") is False
def test_openapi_delete_success(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector.delete()
vector._client.delete_collection.assert_called_once()
def test_openapi_delete_propagates_errors(monkeypatch):
_install_openapi_stubs(monkeypatch)
vector = AnalyticdbVectorOpenAPI.__new__(AnalyticdbVectorOpenAPI)
vector._collection_name = "collection_1"
vector.config = _config()
vector._client = MagicMock()
vector._client.delete_collection.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):
vector.delete()
@@ -1,427 +0,0 @@
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock
import psycopg2.errors
import pytest
import core.rag.datasource.vdb.analyticdb.analyticdb_vector_sql as sql_module
from core.rag.datasource.vdb.analyticdb.analyticdb_vector_sql import (
AnalyticdbVectorBySql,
AnalyticdbVectorBySqlConfig,
)
from core.rag.models.document import Document
def _config_values() -> dict:
return {
"host": "localhost",
"port": 5432,
"account": "account",
"account_password": "password",
"min_connection": 1,
"max_connection": 2,
"namespace": "dify",
}
@pytest.mark.parametrize(
("field", "value", "error_message"),
[
("host", "", "ANALYTICDB_HOST"),
("port", 0, "ANALYTICDB_PORT"),
("account", "", "ANALYTICDB_ACCOUNT"),
("account_password", "", "ANALYTICDB_PASSWORD"),
("min_connection", 0, "ANALYTICDB_MIN_CONNECTION"),
("max_connection", 0, "ANALYTICDB_MAX_CONNECTION"),
],
)
def test_sql_config_required_fields(field, value, error_message):
values = _config_values()
values[field] = value
with pytest.raises(ValueError, match=error_message):
AnalyticdbVectorBySqlConfig.model_validate(values)
def test_sql_config_rejects_min_connection_greater_than_max_connection():
values = _config_values()
values["min_connection"] = 10
values["max_connection"] = 2
with pytest.raises(ValueError, match="ANALYTICDB_MIN_CONNECTION should less than ANALYTICDB_MAX_CONNECTION"):
AnalyticdbVectorBySqlConfig.model_validate(values)
def test_initialize_skips_when_cache_exists(monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(sql_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(sql_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(sql_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector._initialize_vector_database = MagicMock()
vector._initialize()
vector._initialize_vector_database.assert_not_called()
def test_initialize_runs_when_cache_is_missing(monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(sql_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(sql_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(sql_module.redis_client, "set", MagicMock())
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector._initialize_vector_database = MagicMock()
vector._initialize()
vector._initialize_vector_database.assert_called_once()
sql_module.redis_client.set.assert_called_once()
def test_create_connection_pool_uses_psycopg2_pool(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector.databaseName = "knowledgebase"
pool_instance = MagicMock()
monkeypatch.setattr(sql_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool_instance))
pool = vector._create_connection_pool()
assert pool is pool_instance
sql_module.psycopg2.pool.SimpleConnectionPool.assert_called_once()
def test_get_cursor_context_manager_handles_connection_lifecycle():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
cursor = MagicMock()
connection = MagicMock()
connection.cursor.return_value = cursor
pool = MagicMock()
pool.getconn.return_value = connection
vector.pool = pool
with vector._get_cursor() as cur:
assert cur is cursor
cursor.close.assert_called_once()
connection.commit.assert_called_once()
pool.putconn.assert_called_once_with(connection)
def test_add_texts_inserts_only_documents_with_metadata(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
monkeypatch.setattr(sql_module.uuid, "uuid4", lambda: "prefix-id")
monkeypatch.setattr(sql_module.psycopg2.extras, "execute_batch", MagicMock())
docs = [
Document(page_content="doc 1", metadata={"doc_id": "d1", "document_id": "doc-1"}),
SimpleNamespace(page_content="doc 2", metadata=None),
]
vector.add_texts(docs, [[0.1, 0.2], [0.2, 0.3]])
execute_args = sql_module.psycopg2.extras.execute_batch.call_args.args
assert execute_args[0] is cursor
assert len(execute_args[2]) == 1
def test_text_exists_returns_true_and_false_based_on_query_result():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
cursor.fetchone.return_value = ("row",)
assert vector.text_exists("d1") is True
cursor.fetchone.return_value = None
assert vector.text_exists("d1") is False
def test_delete_by_ids_handles_empty_input_and_missing_table_error():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
vector.delete_by_ids([])
cursor.execute.assert_not_called()
cursor.execute.side_effect = psycopg2.errors.UndefinedTable("relation does not exist")
vector.delete_by_ids(["d1"])
def test_delete_by_metadata_field_handles_missing_table_error():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
cursor.execute.side_effect = psycopg2.errors.UndefinedTable("relation does not exist")
vector.delete_by_metadata_field("document_id", "doc-1")
@pytest.mark.parametrize("invalid_top_k", [0, "x", -1])
def test_search_by_vector_validates_top_k(invalid_top_k):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_vector([0.1, 0.2], top_k=invalid_top_k)
def test_search_by_vector_returns_documents_above_threshold():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
cursor.__iter__.return_value = iter(
[
("id1", [1.0], 0.8, "content 1", {"doc_id": "id1", "document_id": "doc-1"}),
("id2", [2.0], 0.3, "content 2", {"doc_id": "id2", "document_id": "doc-2"}),
]
)
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["doc-1"])
assert len(docs) == 1
assert docs[0].page_content == "content 1"
assert docs[0].metadata["score"] == 0.8
@pytest.mark.parametrize("invalid_top_k", [0, "x", -1])
def test_search_by_full_text_validates_top_k(invalid_top_k):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_full_text("query", top_k=invalid_top_k)
def test_search_by_full_text_returns_documents():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
cursor.__iter__.return_value = iter(
[
("id1", [1.0], "content 1", {"doc_id": "id1", "document_id": "doc-1"}, 0.9),
]
)
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
docs = vector.search_by_full_text("query", top_k=1, document_ids_filter=["doc-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == 0.9
assert docs[0].page_content == "content 1"
def test_delete_drops_table():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
vector.delete()
cursor.execute.assert_called_once()
def test_init_normalizes_collection_name_and_creates_pool_when_missing(monkeypatch):
config = AnalyticdbVectorBySqlConfig(**_config_values())
created_pool = MagicMock()
monkeypatch.setattr(AnalyticdbVectorBySql, "_initialize", MagicMock())
monkeypatch.setattr(AnalyticdbVectorBySql, "_create_connection_pool", MagicMock(return_value=created_pool))
vector = AnalyticdbVectorBySql("My_Collection", config)
assert vector._collection_name == "my_collection"
assert vector.table_name == "dify.my_collection"
assert vector.databaseName == "knowledgebase"
assert vector.pool is created_pool
def test_initialize_vector_database_handles_existing_database_and_search_config(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector.databaseName = "knowledgebase"
bootstrap_cursor = MagicMock()
bootstrap_connection = MagicMock()
bootstrap_connection.cursor.return_value = bootstrap_cursor
bootstrap_cursor.execute.side_effect = RuntimeError("database already exists")
monkeypatch.setattr(sql_module.psycopg2, "connect", MagicMock(return_value=bootstrap_connection))
worker_cursor = MagicMock()
worker_connection = MagicMock()
worker_cursor.connection = worker_connection
def _execute(sql, *args, **kwargs):
if "CREATE TEXT SEARCH CONFIGURATION zh_cn" in sql:
raise RuntimeError("already exists")
worker_cursor.execute.side_effect = _execute
pooled_connection = MagicMock()
pooled_connection.cursor.return_value = worker_cursor
pool = MagicMock()
pool.getconn.return_value = pooled_connection
vector._create_connection_pool = MagicMock(return_value=pool)
vector._initialize_vector_database()
bootstrap_cursor.close.assert_called_once()
bootstrap_connection.close.assert_called_once()
vector._create_connection_pool.assert_called_once()
assert any(
"CREATE OR REPLACE FUNCTION public.to_tsquery_from_text" in call.args[0]
for call in worker_cursor.execute.call_args_list
)
assert any("CREATE SCHEMA IF NOT EXISTS dify" in call.args[0] for call in worker_cursor.execute.call_args_list)
def test_initialize_vector_database_raises_runtime_error_when_zhparser_fails(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector.databaseName = "knowledgebase"
bootstrap_cursor = MagicMock()
bootstrap_connection = MagicMock()
bootstrap_connection.cursor.return_value = bootstrap_cursor
monkeypatch.setattr(sql_module.psycopg2, "connect", MagicMock(return_value=bootstrap_connection))
worker_cursor = MagicMock()
worker_connection = MagicMock()
worker_cursor.connection = worker_connection
worker_cursor.execute.side_effect = RuntimeError("zhparser unavailable")
pooled_connection = MagicMock()
pooled_connection.cursor.return_value = worker_cursor
pool = MagicMock()
pool.getconn.return_value = pooled_connection
vector._create_connection_pool = MagicMock(return_value=pool)
with pytest.raises(RuntimeError, match="Failed to create zhparser extension"):
vector._initialize_vector_database()
worker_connection.rollback.assert_called_once()
def test_create_collection_if_not_exists_creates_table_indexes_and_cache(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector._collection_name = "collection"
vector.table_name = "dify.collection"
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(sql_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(sql_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(sql_module.redis_client, "set", MagicMock())
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
vector.create_collection_if_not_exists(embedding_dimension=3)
assert any("CREATE TABLE IF NOT EXISTS dify.collection" in call.args[0] for call in cursor.execute.call_args_list)
assert any("CREATE INDEX collection_embedding_idx" in call.args[0] for call in cursor.execute.call_args_list)
sql_module.redis_client.set.assert_called_once()
def test_create_collection_if_not_exists_raises_for_non_existing_error(monkeypatch):
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.config = AnalyticdbVectorBySqlConfig(**_config_values())
vector._collection_name = "collection"
vector.table_name = "dify.collection"
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(sql_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(sql_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(sql_module.redis_client, "set", MagicMock())
cursor = MagicMock()
cursor.execute.side_effect = RuntimeError("permission denied")
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
with pytest.raises(RuntimeError, match="permission denied"):
vector.create_collection_if_not_exists(embedding_dimension=3)
def test_delete_methods_raise_when_error_is_not_missing_table():
vector = AnalyticdbVectorBySql.__new__(AnalyticdbVectorBySql)
vector.table_name = "dify.collection"
cursor = MagicMock()
@contextmanager
def _cursor_context():
yield cursor
vector._get_cursor = _cursor_context
cursor.execute.side_effect = RuntimeError("unexpected delete failure")
with pytest.raises(RuntimeError, match="unexpected delete failure"):
vector.delete_by_ids(["doc-1"])
cursor.execute.side_effect = RuntimeError("unexpected metadata failure")
with pytest.raises(RuntimeError, match="unexpected metadata failure"):
vector.delete_by_metadata_field("document_id", "doc-1")
@@ -1,551 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_pymochow_modules():
pymochow = types.ModuleType("pymochow")
pymochow.__path__ = []
pymochow_auth = types.ModuleType("pymochow.auth")
pymochow_auth.__path__ = []
pymochow_credentials = types.ModuleType("pymochow.auth.bce_credentials")
pymochow_configuration = types.ModuleType("pymochow.configuration")
pymochow_exception = types.ModuleType("pymochow.exception")
pymochow_model = types.ModuleType("pymochow.model")
pymochow_model.__path__ = []
pymochow_model_database = types.ModuleType("pymochow.model.database")
pymochow_model_enum = types.ModuleType("pymochow.model.enum")
pymochow_model_schema = types.ModuleType("pymochow.model.schema")
pymochow_model_table = types.ModuleType("pymochow.model.table")
class _SimpleObject:
def __init__(self, *args, **kwargs):
self.args = args
for key, value in kwargs.items():
setattr(self, key, value)
class ServerError(Exception):
def __init__(self, code):
super().__init__(f"server error {code}")
self.code = code
class ServerErrCode:
TABLE_NOT_EXIST = 1001
DB_ALREADY_EXIST = 1002
class IndexType:
__members__ = {"HNSW": "HNSW"}
class MetricType:
__members__ = {"IP": "IP"}
class IndexState:
NORMAL = "NORMAL"
class TableState:
NORMAL = "NORMAL"
class InvertedIndexAnalyzer:
DEFAULT_ANALYZER = "DEFAULT_ANALYZER"
class InvertedIndexParseMode:
COARSE_MODE = "COARSE_MODE"
class InvertedIndexFieldAttribute:
ANALYZED = "ANALYZED"
class FieldType:
STRING = "STRING"
TEXT = "TEXT"
JSON = "JSON"
FLOAT_VECTOR = "FLOAT_VECTOR"
pymochow.MochowClient = _SimpleObject
pymochow_credentials.BceCredentials = _SimpleObject
pymochow_configuration.Configuration = _SimpleObject
pymochow_exception.ServerError = ServerError
pymochow_model_database.Database = _SimpleObject
pymochow_model_enum.FieldType = FieldType
pymochow_model_enum.IndexState = IndexState
pymochow_model_enum.IndexType = IndexType
pymochow_model_enum.MetricType = MetricType
pymochow_model_enum.ServerErrCode = ServerErrCode
pymochow_model_enum.TableState = TableState
for cls_name in [
"AutoBuildRowCountIncrement",
"Field",
"FilteringIndex",
"HNSWParams",
"InvertedIndex",
"InvertedIndexParams",
"Schema",
"VectorIndex",
]:
setattr(pymochow_model_schema, cls_name, _SimpleObject)
pymochow_model_schema.InvertedIndexAnalyzer = InvertedIndexAnalyzer
pymochow_model_schema.InvertedIndexFieldAttribute = InvertedIndexFieldAttribute
pymochow_model_schema.InvertedIndexParseMode = InvertedIndexParseMode
for cls_name in ["AnnSearch", "BM25SearchRequest", "HNSWSearchParams", "Partition", "Row"]:
setattr(pymochow_model_table, cls_name, _SimpleObject)
pymochow.auth = pymochow_auth
pymochow.model = pymochow_model
pymochow_auth.bce_credentials = pymochow_credentials
pymochow_model.database = pymochow_model_database
pymochow_model.enum = pymochow_model_enum
pymochow_model.schema = pymochow_model_schema
pymochow_model.table = pymochow_model_table
modules = {
"pymochow": pymochow,
"pymochow.auth": pymochow_auth,
"pymochow.auth.bce_credentials": pymochow_credentials,
"pymochow.configuration": pymochow_configuration,
"pymochow.exception": pymochow_exception,
"pymochow.model": pymochow_model,
"pymochow.model.database": pymochow_model_database,
"pymochow.model.enum": pymochow_model_enum,
"pymochow.model.schema": pymochow_model_schema,
"pymochow.model.table": pymochow_model_table,
}
return modules
@pytest.fixture
def baidu_module(monkeypatch):
for name, module in _build_fake_pymochow_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.baidu.baidu_vector as module
return importlib.reload(module)
def test_baidu_config_validation(baidu_module):
values = {
"endpoint": "https://example.com",
"account": "account",
"api_key": "key",
"database": "database",
}
config = baidu_module.BaiduConfig.model_validate(values)
assert config.endpoint == "https://example.com"
for key, error_message in [
("endpoint", "BAIDU_VECTOR_DB_ENDPOINT"),
("account", "BAIDU_VECTOR_DB_ACCOUNT"),
("api_key", "BAIDU_VECTOR_DB_API_KEY"),
("database", "BAIDU_VECTOR_DB_DATABASE"),
]:
invalid = dict(values)
invalid[key] = ""
with pytest.raises(ValueError, match=error_message):
baidu_module.BaiduConfig.model_validate(invalid)
def test_get_search_result_handles_metadata_and_threshold(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
response = SimpleNamespace(
rows=[
{"row": {"page_content": "doc1", "metadata": '{"document_id":"d1"}'}, "score": 0.9},
{"row": {"page_content": "doc2", "metadata": {"document_id": "d2"}}, "score": 0.4},
{"row": {"page_content": "doc3", "metadata": 123}, "score": 0.95},
]
)
docs = vector._get_search_res(response, score_threshold=0.8)
assert len(docs) == 2
assert docs[0].page_content == "doc1"
assert docs[0].metadata["score"] == 0.9
assert docs[1].page_content == "doc3"
def test_delete_by_ids_and_delete_by_metadata_field(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
table = MagicMock()
vector._db = MagicMock()
vector._db.table.return_value = table
vector._collection_name = "collection_1"
vector.delete_by_ids([])
table.delete.assert_not_called()
vector.delete_by_ids(["id1", "id2"])
table.delete.assert_called_once()
table.delete.reset_mock()
vector.delete_by_metadata_field("source", 'abc"def')
delete_filter = table.delete.call_args.kwargs["filter"]
assert delete_filter == 'metadata["source"] = "abc\\"def"'
def test_delete_handles_table_not_exist_error_and_raises_for_other_codes(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._db = MagicMock()
vector._db.drop_table.side_effect = baidu_module.ServerError(baidu_module.ServerErrCode.TABLE_NOT_EXIST)
vector.delete()
vector._db.drop_table.side_effect = baidu_module.ServerError(9999)
with pytest.raises(baidu_module.ServerError):
vector.delete()
def test_init_database_uses_existing_or_creates_when_missing(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._client = MagicMock()
vector._client_config = SimpleNamespace(database="my_db")
vector._client.list_databases.return_value = [SimpleNamespace(database_name="my_db")]
vector._client.database.return_value = "existing_db"
assert vector._init_database() == "existing_db"
vector._client.list_databases.return_value = []
vector._client.database.return_value = "created_db"
vector._client.create_database.side_effect = None
assert vector._init_database() == "created_db"
vector._client.create_database.side_effect = baidu_module.ServerError(baidu_module.ServerErrCode.DB_ALREADY_EXIST)
assert vector._init_database() == "created_db"
def test_table_existed_checks_table_access(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._db = MagicMock()
vector._db.table.return_value = MagicMock()
assert vector._table_existed() is True
vector._db.table.side_effect = baidu_module.ServerError(baidu_module.ServerErrCode.TABLE_NOT_EXIST)
assert vector._table_existed() is False
vector._db.table.side_effect = baidu_module.ServerError(9999)
with pytest.raises(baidu_module.ServerError):
vector._table_existed()
def test_search_methods_delegate_to_database_table(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._db = MagicMock()
vector._get_search_res = MagicMock(return_value=[Document(page_content="doc", metadata={"doc_id": "1"})])
table = MagicMock()
vector._db.table.return_value = table
table.search.return_value = "vector_result"
table.bm25_search.return_value = "bm25_result"
result1 = vector.search_by_vector([0.1, 0.2], top_k=3, document_ids_filter=["doc-1"], score_threshold=0.2)
result2 = vector.search_by_full_text("query", top_k=3, document_ids_filter=["doc-1"], score_threshold=0.2)
assert result1 == vector._get_search_res.return_value
assert result2 == vector._get_search_res.return_value
assert vector._get_search_res.call_count == 2
def test_factory_initializes_collection_name_and_index_struct(baidu_module, monkeypatch):
factory = baidu_module.BaiduVectorFactory()
dataset = SimpleNamespace(id="dataset-1", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(baidu_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_ENDPOINT", "https://endpoint")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_CONNECTION_TIMEOUT_MS", 1000)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_ACCOUNT", "account")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_API_KEY", "key")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_DATABASE", "database")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_SHARD", 1)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_REPLICAS", 1)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_INVERTED_INDEX_ANALYZER", "DEFAULT_ANALYZER")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_INVERTED_INDEX_PARSER_MODE", "COARSE_MODE")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT", 500)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT_RATIO", 0.05)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_REBUILD_INDEX_TIMEOUT_IN_SECONDS", 300)
with patch.object(baidu_module, "BaiduVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
assert vector_cls.call_args.kwargs["collection_name"] == "auto_collection"
assert dataset.index_struct is not None
def test_init_get_type_to_index_struct_and_create_delegate(baidu_module, monkeypatch):
init_client = MagicMock(return_value="client")
init_database = MagicMock(return_value="database")
monkeypatch.setattr(baidu_module.BaiduVector, "_init_client", init_client)
monkeypatch.setattr(baidu_module.BaiduVector, "_init_database", init_database)
config = baidu_module.BaiduConfig(
endpoint="https://example.com",
account="account",
api_key="key",
database="db",
)
vector = baidu_module.BaiduVector(collection_name="my_collection", config=config)
assert vector.get_type() == baidu_module.VectorType.BAIDU
assert vector.to_index_struct()["vector_store"]["class_prefix"] == "my_collection"
assert vector._client == "client"
assert vector._db == "database"
vector._create_table = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="p1", metadata={"doc_id": "d1"})]
vector.create(docs, [[0.1, 0.2]])
vector._create_table.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_batches_rows(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
table = MagicMock()
vector._db = MagicMock()
vector._db.table.return_value = table
docs = [
Document(page_content="doc-1", metadata={"doc_id": "id-1", "document_id": "doc-1"}),
Document(page_content="doc-2", metadata={"doc_id": "id-2", "document_id": "doc-2"}),
]
vector.add_texts(docs, [[0.1, 0.2], [0.3, 0.4]])
assert table.upsert.call_count == 1
inserted_rows = table.upsert.call_args.kwargs["rows"]
assert len(inserted_rows) == 2
def test_add_texts_batches_more_than_batch_size(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
table = MagicMock()
vector._db = MagicMock()
vector._db.table.return_value = table
docs = [
Document(page_content=f"doc-{idx}", metadata={"doc_id": f"id-{idx}", "document_id": f"doc-{idx}"})
for idx in range(1001)
]
embeddings = [[0.1, 0.2] for _ in range(1001)]
vector.add_texts(docs, embeddings)
assert table.upsert.call_count == 2
assert len(table.upsert.call_args_list[0].kwargs["rows"]) == 1000
assert len(table.upsert.call_args_list[1].kwargs["rows"]) == 1
def test_text_exists_returns_false_when_query_code_is_not_success(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
table = MagicMock()
vector._db = MagicMock()
vector._db.table.return_value = table
table.query.return_value = SimpleNamespace(code=0)
assert vector.text_exists("id-1") is True
table.query.return_value = SimpleNamespace(code=1)
assert vector.text_exists("id-1") is False
table.query.return_value = None
assert vector.text_exists("id-1") is False
def test_get_search_result_handles_invalid_metadata_json(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
response = SimpleNamespace(rows=[{"row": {"page_content": "doc1", "metadata": "{bad json"}, "score": 0.7}])
docs = vector._get_search_res(response, score_threshold=0.1)
assert len(docs) == 1
assert docs[0].metadata["score"] == 0.7
assert "document_id" not in docs[0].metadata
def test_init_client_constructs_configuration_and_client(baidu_module, monkeypatch):
credentials = MagicMock(return_value="credentials")
configuration = MagicMock(return_value="configuration")
client_cls = MagicMock(return_value="client")
monkeypatch.setattr(baidu_module, "BceCredentials", credentials)
monkeypatch.setattr(baidu_module, "Configuration", configuration)
monkeypatch.setattr(baidu_module, "MochowClient", client_cls)
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
config = SimpleNamespace(
account="account",
api_key="key",
endpoint="https://endpoint",
connection_timeout_in_mills=12_345,
)
client = vector._init_client(config)
assert client == "client"
credentials.assert_called_once_with("account", "key")
configuration.assert_called_once_with(
credentials="credentials",
endpoint="https://endpoint",
connection_timeout_in_mills=12_345,
)
client_cls.assert_called_once_with("configuration")
def test_init_database_raises_for_unknown_create_database_error(baidu_module):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._client = MagicMock()
vector._client_config = SimpleNamespace(database="my_db")
vector._client.list_databases.return_value = []
vector._client.create_database.side_effect = baidu_module.ServerError(9999)
with pytest.raises(baidu_module.ServerError):
vector._init_database()
def test_create_table_handles_cache_and_validation_paths(baidu_module, monkeypatch):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._client_config = SimpleNamespace(
index_type="HNSW",
metric_type="IP",
inverted_index_analyzer="DEFAULT_ANALYZER",
inverted_index_parser_mode="COARSE_MODE",
auto_build_row_count_increment=500,
auto_build_row_count_increment_ratio=0.05,
rebuild_index_timeout_in_seconds=300,
replicas=1,
shard=1,
)
vector._db = MagicMock()
table = MagicMock()
table.state = baidu_module.TableState.NORMAL
vector._db.describe_table.return_value = table
vector._table_existed = MagicMock(return_value=False)
vector.delete = MagicMock()
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(baidu_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(baidu_module.redis_client, "set", MagicMock())
monkeypatch.setattr(baidu_module.time, "sleep", lambda _s: None)
monkeypatch.setattr(vector, "_wait_for_index_ready", MagicMock())
# Cached table skips all work.
monkeypatch.setattr(baidu_module.redis_client, "get", MagicMock(return_value=1))
vector._create_table(3)
vector._db.create_table.assert_not_called()
# Existing table also skips creation.
monkeypatch.setattr(baidu_module.redis_client, "get", MagicMock(return_value=None))
vector._table_existed.return_value = True
vector._create_table(3)
vector._db.create_table.assert_not_called()
# Create table when cache is empty and table does not exist.
vector._table_existed.return_value = False
vector._create_table(3)
vector._db.create_table.assert_called_once()
baidu_module.redis_client.set.assert_called_once_with("vector_indexing_collection_1", 1, ex=3600)
table.rebuild_index.assert_called_once_with(vector.vector_index)
vector._wait_for_index_ready.assert_called_once_with(table, 3600)
def test_create_table_raises_for_invalid_index_or_metric(baidu_module, monkeypatch):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._db = MagicMock()
vector._table_existed = MagicMock(return_value=False)
vector.delete = MagicMock()
vector._client_config = SimpleNamespace(
index_type="INVALID",
metric_type="IP",
inverted_index_analyzer="DEFAULT_ANALYZER",
inverted_index_parser_mode="COARSE_MODE",
auto_build_row_count_increment=500,
auto_build_row_count_increment_ratio=0.05,
rebuild_index_timeout_in_seconds=300,
replicas=1,
shard=1,
)
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(baidu_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(baidu_module.redis_client, "get", MagicMock(return_value=None))
with pytest.raises(ValueError, match="unsupported index_type"):
vector._create_table(3)
vector._client_config.index_type = "HNSW"
vector._client_config.metric_type = "INVALID"
with pytest.raises(ValueError, match="unsupported metric_type"):
vector._create_table(3)
def test_create_table_raises_timeout_if_table_never_becomes_normal(baidu_module, monkeypatch):
vector = baidu_module.BaiduVector.__new__(baidu_module.BaiduVector)
vector._collection_name = "collection_1"
vector._client_config = SimpleNamespace(
index_type="HNSW",
metric_type="IP",
inverted_index_analyzer="DEFAULT_ANALYZER",
inverted_index_parser_mode="COARSE_MODE",
auto_build_row_count_increment=500,
auto_build_row_count_increment_ratio=0.05,
rebuild_index_timeout_in_seconds=300,
replicas=1,
shard=1,
)
vector._db = MagicMock()
vector._db.describe_table.return_value = SimpleNamespace(state="CREATING")
vector._table_existed = MagicMock(return_value=False)
vector.delete = MagicMock()
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(baidu_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(baidu_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(baidu_module.time, "sleep", lambda _s: None)
monkeypatch.setattr(baidu_module.time, "time", MagicMock(side_effect=[0, 301]))
with pytest.raises(TimeoutError, match="Table creation timeout"):
vector._create_table(3)
def test_factory_uses_existing_collection_prefix_when_index_struct_exists(baidu_module, monkeypatch):
factory = baidu_module.BaiduVectorFactory()
dataset = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_ENDPOINT", "https://endpoint")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_CONNECTION_TIMEOUT_MS", 1000)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_ACCOUNT", "account")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_API_KEY", "key")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_DATABASE", "database")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_SHARD", 1)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_REPLICAS", 1)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_INVERTED_INDEX_ANALYZER", "DEFAULT_ANALYZER")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_INVERTED_INDEX_PARSER_MODE", "COARSE_MODE")
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT", 500)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_AUTO_BUILD_ROW_COUNT_INCREMENT_RATIO", 0.05)
monkeypatch.setattr(baidu_module.dify_config, "BAIDU_VECTOR_DB_REBUILD_INDEX_TIMEOUT_IN_SECONDS", 300)
with patch.object(baidu_module, "BaiduVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
assert vector_cls.call_args.kwargs["collection_name"] == "existing_collection"
@@ -1,199 +0,0 @@
import importlib
import sys
import types
from collections import UserDict
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_chroma_modules():
chroma = types.ModuleType("chromadb")
chroma.DEFAULT_TENANT = "default_tenant"
chroma.DEFAULT_DATABASE = "default_database"
class Settings:
def __init__(self, **kwargs):
for key, value in kwargs.items():
setattr(self, key, value)
class QueryResult(UserDict):
pass
class _Collection:
def __init__(self):
self.upsert = MagicMock()
self.delete = MagicMock()
self.query = MagicMock()
self.get = MagicMock(return_value={})
class _Client:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.collection = _Collection()
self.get_or_create_collection = MagicMock(return_value=self.collection)
self.delete_collection = MagicMock()
chroma.Settings = Settings
chroma.QueryResult = QueryResult
chroma.HttpClient = _Client
return chroma
@pytest.fixture
def chroma_module(monkeypatch):
fake_chroma = _build_fake_chroma_modules()
monkeypatch.setitem(sys.modules, "chromadb", fake_chroma)
import core.rag.datasource.vdb.chroma.chroma_vector as module
return importlib.reload(module)
def test_chroma_config_to_params_builds_expected_payload(chroma_module):
config = chroma_module.ChromaConfig(
host="localhost",
port=8000,
tenant="tenant-1",
database="db-1",
auth_provider="provider",
auth_credentials="credentials",
)
params = config.to_chroma_params()
assert params["host"] == "localhost"
assert params["port"] == 8000
assert params["tenant"] == "tenant-1"
assert params["database"] == "db-1"
assert params["ssl"] is False
assert params["settings"].chroma_client_auth_provider == "provider"
assert params["settings"].chroma_client_auth_credentials == "credentials"
def test_create_collection_uses_redis_lock_and_cache(chroma_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(chroma_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(chroma_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(chroma_module.redis_client, "set", MagicMock())
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
vector.create_collection("collection_1")
vector._client.get_or_create_collection.assert_called_once_with("collection_1")
chroma_module.redis_client.set.assert_called_once()
def test_create_with_empty_texts_is_noop(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
vector.create([], [])
vector._client.get_or_create_collection.assert_not_called()
def test_create_with_texts_creates_collection_and_upserts(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
docs = [Document(page_content="hello", metadata={"doc_id": "d1", "document_id": "doc-1"})]
vector.create(docs, [[0.1, 0.2]])
vector._client.get_or_create_collection.assert_called()
vector._client.collection.upsert.assert_called_once()
def test_delete_methods_and_text_exists(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
vector.delete_by_ids([])
vector._client.collection.delete.assert_not_called()
vector.delete_by_ids(["id-1"])
vector._client.collection.delete.assert_called_with(ids=["id-1"])
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.collection.delete.assert_called_with(where={"document_id": {"$eq": "doc-1"}})
vector._client.collection.get.return_value = {"ids": ["id-1"]}
assert vector.text_exists("id-1") is True
vector._client.collection.get.return_value = {}
assert vector.text_exists("id-2") is False
vector.delete()
vector._client.delete_collection.assert_called_once_with("collection_1")
def test_search_by_vector_handles_empty_results(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
vector._client.collection.query.return_value = {"ids": [], "documents": [], "metadatas": [], "distances": []}
assert vector.search_by_vector([0.1, 0.2], top_k=2) == []
def test_search_by_vector_applies_score_threshold_and_sorting(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
vector._client.collection.query.return_value = {
"ids": [["id-1", "id-2"]],
"documents": [["doc high", "doc low"]],
"metadatas": [[{"doc_id": "id-1"}, {"doc_id": "id-2"}]],
"distances": [[0.1, 0.8]],
}
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["doc-1"])
assert len(docs) == 1
assert docs[0].page_content == "doc high"
assert docs[0].metadata["score"] == 0.9
def test_search_by_full_text_returns_empty_list(chroma_module):
vector = chroma_module.ChromaVector(
collection_name="collection_1",
config=chroma_module.ChromaConfig(host="localhost", port=8000, tenant="t", database="d"),
)
assert vector.search_by_full_text("query") == []
def test_factory_init_vector_uses_existing_or_generated_collection(chroma_module, monkeypatch):
factory = chroma_module.ChromaVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1", index_struct_dict={"vector_store": {"class_prefix": "EXISTING"}}, index_struct=None
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(chroma_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_HOST", "localhost")
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_PORT", 8000)
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_TENANT", None)
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_DATABASE", None)
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_AUTH_PROVIDER", None)
monkeypatch.setattr(chroma_module.dify_config, "CHROMA_AUTH_CREDENTIALS", None)
with patch.object(chroma_module, "ChromaVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -1,927 +0,0 @@
import importlib
import queue
import sys
import types
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_clickzetta_module():
clickzetta = types.ModuleType("clickzetta")
class _FakeCursor:
def __init__(self):
self.execute = MagicMock()
self.executemany = MagicMock()
self.fetchall = MagicMock(return_value=[])
self.fetchone = MagicMock(return_value=(0,))
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return False
class _FakeConnection:
def __init__(self):
self.cursor_obj = _FakeCursor()
def cursor(self):
return self.cursor_obj
def close(self):
return None
def connect(**_kwargs):
return _FakeConnection()
clickzetta.connect = connect
return clickzetta
@pytest.fixture
def clickzetta_module(monkeypatch):
monkeypatch.setitem(sys.modules, "clickzetta", _build_fake_clickzetta_module())
import core.rag.datasource.vdb.clickzetta.clickzetta_vector as module
return importlib.reload(module)
def _config(module):
return module.ClickzettaConfig(
username="username",
password="password",
instance="instance",
service="service",
workspace="workspace",
vcluster="cluster",
schema_name="dify",
)
@pytest.mark.parametrize(
("field", "error_message"),
[
("username", "CLICKZETTA_USERNAME"),
("password", "CLICKZETTA_PASSWORD"),
("instance", "CLICKZETTA_INSTANCE"),
("service", "CLICKZETTA_SERVICE"),
("workspace", "CLICKZETTA_WORKSPACE"),
("vcluster", "CLICKZETTA_VCLUSTER"),
("schema_name", "CLICKZETTA_SCHEMA"),
],
)
def test_clickzetta_config_validation(clickzetta_module, field, error_message):
values = _config(clickzetta_module).model_dump()
values[field] = ""
with pytest.raises(ValueError, match=error_message):
clickzetta_module.ClickzettaConfig.model_validate(values)
def test_parse_metadata_handles_valid_double_encoded_and_invalid_json(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
parsed = vector._parse_metadata('{"document_id":"doc-1"}', "row-1")
assert parsed["doc_id"] == "row-1"
assert parsed["document_id"] == "doc-1"
parsed_double = vector._parse_metadata('"{\\"document_id\\": \\"doc-2\\"}"', "row-2")
assert parsed_double["doc_id"] == "row-2"
assert parsed_double["document_id"] == "doc-2"
parsed_fallback = vector._parse_metadata("not-json", "row-3")
assert parsed_fallback["doc_id"] == "row-3"
assert parsed_fallback["document_id"] == "row-3"
def test_safe_doc_id_and_vector_format_helpers(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
assert vector._format_vector_simple([0.1, 0.2, 0.3]) == "0.1,0.2,0.3"
assert vector._safe_doc_id("abc-123_DEF") == "abc-123_DEF"
assert vector._safe_doc_id("ab c;\n") == "abc"
assert len(vector._safe_doc_id("a" * 300)) == 255
def test_table_exists_returns_false_for_not_found_and_other_exceptions(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
@contextmanager
def _ctx_not_found():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.execute.side_effect = RuntimeError("CZLH-42000 table or view not found")
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_not_found
assert vector._table_exists() is False
@contextmanager
def _ctx_other_error():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.execute.side_effect = RuntimeError("permission denied")
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_other_error
assert vector._table_exists() is False
def test_text_exists_handles_missing_table_and_existing_rows(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=False)
assert vector.text_exists("doc-1") is False
vector._table_exists = MagicMock(return_value=True)
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.fetchone.return_value = (1,)
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
assert vector.text_exists("doc-1") is True
def test_delete_by_ids_and_delete_by_metadata_field_short_circuit(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._execute_write = MagicMock()
vector.delete_by_ids([])
vector._execute_write.assert_not_called()
vector._table_exists = MagicMock(return_value=False)
vector.delete_by_ids(["doc-1"])
vector._execute_write.assert_not_called()
vector.delete_by_metadata_field("document_id", "doc-1")
vector._execute_write.assert_not_called()
def test_search_short_circuit_behaviors(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=False)
assert vector.search_by_vector([0.1, 0.2], top_k=2) == []
vector._config.enable_inverted_index = False
assert vector.search_by_full_text("query", top_k=2) == []
def test_search_by_like_returns_documents_with_default_score(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=True)
vector._parse_metadata = MagicMock(return_value={"document_id": "doc-1", "doc_id": "seg-1"})
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.fetchall.return_value = [("seg-1", "content", '{"document_id":"doc-1"}')]
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
docs = vector._search_by_like("query", top_k=3, document_ids_filter=["doc-1"])
assert len(docs) == 1
assert docs[0].page_content == "content"
assert docs[0].metadata["score"] == 0.5
def test_factory_initializes_clickzetta_vector(clickzetta_module, monkeypatch):
factory = clickzetta_module.ClickzettaVectorFactory()
dataset = SimpleNamespace(id="dataset-1")
monkeypatch.setattr(clickzetta_module.Dataset, "gen_collection_name_by_id", lambda _id: "COLLECTION")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_USERNAME", "username")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_PASSWORD", "password")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_INSTANCE", "instance")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_SERVICE", "service")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_WORKSPACE", "workspace")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_VCLUSTER", "cluster")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_SCHEMA", "dify")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_BATCH_SIZE", 10)
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_ENABLE_INVERTED_INDEX", True)
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_ANALYZER_TYPE", "chinese")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_ANALYZER_MODE", "smart")
monkeypatch.setattr(clickzetta_module.dify_config, "CLICKZETTA_VECTOR_DISTANCE_FUNCTION", "cosine_distance")
with patch.object(clickzetta_module, "ClickzettaVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
assert vector_cls.call_args.kwargs["collection_name"] == "collection"
def test_connection_pool_singleton_and_config_key(clickzetta_module, monkeypatch):
clickzetta_module.ClickzettaConnectionPool._instance = None
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool_1 = clickzetta_module.ClickzettaConnectionPool.get_instance()
pool_2 = clickzetta_module.ClickzettaConnectionPool.get_instance()
key = pool_1._get_config_key(_config(clickzetta_module))
assert pool_1 is pool_2
assert "username:instance:service:workspace:cluster:dify" in key
def test_connection_pool_create_connection_retries_and_configures(clickzetta_module, monkeypatch):
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool = clickzetta_module.ClickzettaConnectionPool()
config = _config(clickzetta_module)
connection = MagicMock()
monkeypatch.setattr(clickzetta_module.time, "sleep", lambda _s: None)
monkeypatch.setattr(
clickzetta_module.clickzetta, "connect", MagicMock(side_effect=[RuntimeError("boom"), connection])
)
pool._configure_connection = MagicMock()
created = pool._create_connection(config)
assert created is connection
assert clickzetta_module.clickzetta.connect.call_count == 2
pool._configure_connection.assert_called_once_with(connection)
def test_connection_pool_create_connection_raises_after_retries(clickzetta_module, monkeypatch):
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool = clickzetta_module.ClickzettaConnectionPool()
config = _config(clickzetta_module)
monkeypatch.setattr(clickzetta_module.time, "sleep", lambda _s: None)
monkeypatch.setattr(clickzetta_module.clickzetta, "connect", MagicMock(side_effect=RuntimeError("boom")))
with pytest.raises(RuntimeError, match="boom"):
pool._create_connection(config)
def test_connection_pool_configure_and_validate_connection(clickzetta_module):
monkeypatch = pytest.MonkeyPatch()
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool = clickzetta_module.ClickzettaConnectionPool()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection = MagicMock()
connection.cursor.return_value = cursor
pool._configure_connection(connection)
assert cursor.execute.call_count >= 2
assert pool._is_connection_valid(connection) is True
bad_connection = MagicMock()
bad_connection.cursor.side_effect = RuntimeError("bad connection")
assert pool._is_connection_valid(bad_connection) is False
monkeypatch.undo()
def test_connection_pool_configure_connection_swallows_errors(clickzetta_module):
monkeypatch = pytest.MonkeyPatch()
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool = clickzetta_module.ClickzettaConnectionPool()
connection = MagicMock()
connection.cursor.side_effect = RuntimeError("cannot configure")
pool._configure_connection(connection)
monkeypatch.undo()
def test_connection_pool_get_return_cleanup_and_shutdown(clickzetta_module, monkeypatch):
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "_start_cleanup_thread", MagicMock())
pool = clickzetta_module.ClickzettaConnectionPool()
config = _config(clickzetta_module)
key = pool._get_config_key(config)
created_connection = MagicMock()
pool._create_connection = MagicMock(return_value=created_connection)
first = pool.get_connection(config)
assert first is created_connection
reusable_connection = MagicMock()
pool._pools[key] = [(reusable_connection, clickzetta_module.time.time())]
pool._is_connection_valid = MagicMock(return_value=True)
reused = pool.get_connection(config)
assert reused is reusable_connection
expired_connection = MagicMock()
pool._pools[key] = [(expired_connection, 0.0)]
pool._is_connection_valid = MagicMock(return_value=False)
monkeypatch.setattr(clickzetta_module.time, "time", MagicMock(return_value=1000.0))
pool.get_connection(config)
expired_connection.close.assert_called_once()
random_connection = MagicMock()
pool._is_connection_valid = MagicMock(return_value=True)
pool.return_connection(config, random_connection)
assert len(pool._pools[key]) == 1
pool._pools[key] = [(MagicMock(), 0.0), (MagicMock(), 1000.0)]
pool._connection_timeout = 10
pool._cleanup_expired_connections()
assert len(pool._pools[key]) == 1
unknown_pool = MagicMock()
pool.return_connection(_config(clickzetta_module).model_copy(update={"workspace": "other"}), unknown_pool)
unknown_pool.close.assert_called_once()
pool.shutdown()
assert pool._shutdown is True
def test_connection_pool_start_cleanup_thread_runs_worker_once(clickzetta_module, monkeypatch):
pool = clickzetta_module.ClickzettaConnectionPool.__new__(clickzetta_module.ClickzettaConnectionPool)
pool._shutdown = False
pool._cleanup_expired_connections = MagicMock(side_effect=lambda: setattr(pool, "_shutdown", True))
monkeypatch.setattr(clickzetta_module.time, "sleep", lambda _s: None)
class _Thread:
def __init__(self, target, daemon):
self._target = target
self.daemon = daemon
self.started = False
def start(self):
self.started = True
self._target()
monkeypatch.setattr(clickzetta_module.threading, "Thread", _Thread)
pool._start_cleanup_thread()
assert pool._cleanup_thread.started is True
pool._cleanup_expired_connections.assert_called_once()
def test_vector_init_connection_context_and_helpers(clickzetta_module, monkeypatch):
pool = MagicMock()
pool.get_connection.return_value = "conn"
monkeypatch.setattr(clickzetta_module.ClickzettaConnectionPool, "get_instance", MagicMock(return_value=pool))
monkeypatch.setattr(clickzetta_module.ClickzettaVector, "_init_write_queue", MagicMock())
vector = clickzetta_module.ClickzettaVector("My-Collection", _config(clickzetta_module))
assert vector._table_name == "my_collection"
assert vector._get_connection() == "conn"
vector._return_connection("conn")
pool.return_connection.assert_called_with(vector._config, "conn")
with vector.get_connection_context() as conn:
assert conn == "conn"
assert pool.return_connection.call_count >= 2
assert vector.get_type() == "clickzetta"
assert vector._ensure_connection() == "conn"
def test_write_queue_initialization_worker_and_execute_write(clickzetta_module, monkeypatch):
class _Thread:
def __init__(self, target, daemon):
self.target = target
self.daemon = daemon
self.started = 0
def start(self):
self.started += 1
monkeypatch.setattr(clickzetta_module.threading, "Thread", _Thread)
clickzetta_module.ClickzettaVector._write_queue = None
clickzetta_module.ClickzettaVector._write_thread = None
clickzetta_module.ClickzettaVector._shutdown = False
clickzetta_module.ClickzettaVector._init_write_queue()
clickzetta_module.ClickzettaVector._init_write_queue()
assert clickzetta_module.ClickzettaVector._write_thread.started == 1
result_queue_ok = queue.Queue()
result_queue_fail = queue.Queue()
clickzetta_module.ClickzettaVector._write_queue = queue.Queue()
clickzetta_module.ClickzettaVector._shutdown = False
clickzetta_module.ClickzettaVector._write_queue.put((lambda x: x + 1, (1,), {}, result_queue_ok))
clickzetta_module.ClickzettaVector._write_queue.put(
(lambda: (_ for _ in ()).throw(RuntimeError("worker error")), (), {}, result_queue_fail)
)
clickzetta_module.ClickzettaVector._write_queue.put(None)
clickzetta_module.ClickzettaVector._write_worker()
assert result_queue_ok.get() == (True, 2)
failed = result_queue_fail.get()
assert failed[0] is False
assert isinstance(failed[1], RuntimeError)
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
clickzetta_module.ClickzettaVector._write_queue = None
with pytest.raises(RuntimeError, match="Write queue not initialized"):
vector._execute_write(lambda: None)
class _ImmediateSuccessQueue:
def put(self, task):
func, args, kwargs, result_q = task
result_q.put((True, func(*args, **kwargs)))
clickzetta_module.ClickzettaVector._write_queue = _ImmediateSuccessQueue()
assert vector._execute_write(lambda x: x * 2, 3) == 6
class _ImmediateFailQueue:
def put(self, task):
_, _, _, result_q = task
result_q.put((False, ValueError("write failed")))
clickzetta_module.ClickzettaVector._write_queue = _ImmediateFailQueue()
with pytest.raises(ValueError, match="write failed"):
vector._execute_write(lambda: None)
def test_table_exists_true_and_create_invokes_write_and_add_texts(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
@contextmanager
def _ctx_exists():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_exists
assert vector._table_exists() is True
vector._execute_write = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="content", metadata={"doc_id": "d1"})]
vector.create(docs, [[0.1, 0.2]])
vector._execute_write.assert_called_once()
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_create_table_and_indexes_paths(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._create_vector_index = MagicMock()
vector._create_inverted_index = MagicMock()
vector._table_exists = MagicMock(return_value=True)
vector._create_table_and_indexes([[0.1, 0.2]])
vector._create_vector_index.assert_not_called()
vector._table_exists = MagicMock(return_value=False)
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
vector._create_table_and_indexes([[0.1, 0.2, 0.3]])
vector._create_vector_index.assert_called_once()
vector._create_inverted_index.assert_called_once()
vector._config.enable_inverted_index = False
vector._create_vector_index.reset_mock()
vector._create_inverted_index.reset_mock()
vector._create_table_and_indexes([])
vector._create_vector_index.assert_called_once()
vector._create_inverted_index.assert_not_called()
def test_create_vector_index_branches(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
cursor = MagicMock()
cursor.fetchall.return_value = [("idx_table_vector", "embedding_vector")]
vector._create_vector_index(cursor)
assert cursor.execute.call_count == 1
cursor.reset_mock()
cursor.execute.side_effect = [RuntimeError("show index failed"), None]
vector._create_vector_index(cursor)
assert cursor.execute.call_count == 2
cursor.reset_mock()
cursor.execute.side_effect = [None, RuntimeError("already exists")]
cursor.fetchall.return_value = []
vector._create_vector_index(cursor)
cursor.reset_mock()
cursor.execute.side_effect = [None, RuntimeError("unexpected")]
cursor.fetchall.return_value = []
with pytest.raises(RuntimeError, match="unexpected"):
vector._create_vector_index(cursor)
def test_create_inverted_index_branches(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
cursor = MagicMock()
cursor.fetchall.return_value = [("idx_table_1_text", "INVERTED", "page_content")]
vector._create_inverted_index(cursor)
assert cursor.execute.call_count == 1
cursor.reset_mock()
cursor.execute.side_effect = [RuntimeError("show failed"), None]
vector._create_inverted_index(cursor)
assert cursor.execute.call_count == 2
cursor.reset_mock()
cursor.execute.side_effect = [
None,
RuntimeError("already has index"),
None,
]
cursor.fetchall.return_value = [("idx_table_1_text", "INVERTED", "page_content")]
vector._create_inverted_index(cursor)
cursor.reset_mock()
cursor.execute.side_effect = [None, RuntimeError("other create failure")]
cursor.fetchall.return_value = []
vector._create_inverted_index(cursor)
def test_add_texts_batches_and_insert_batch_behaviors(clickzetta_module, monkeypatch):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._config.batch_size = 2
vector._table_name = "table_1"
vector._execute_write = MagicMock()
vector._safe_doc_id = MagicMock(side_effect=lambda doc_id: str(doc_id))
docs = [
Document(page_content="doc-1", metadata={"doc_id": "id-1"}),
Document(page_content="doc-2", metadata={"doc_id": "id-2"}),
Document(page_content="doc-3", metadata={"doc_id": "id-3"}),
]
vectors = [[0.1], [0.2], [0.3]]
vector.add_texts([], [])
vector._execute_write.assert_not_called()
added_ids = vector.add_texts(docs, vectors)
assert added_ids == ["id-1", "id-2", "id-3"]
assert vector._execute_write.call_count == 2
assert vector._execute_write.call_args_list[0].args == (
vector._insert_batch,
docs[:2],
vectors[:2],
["id-1", "id-2"],
0,
2,
2,
)
assert vector._execute_write.call_args_list[1].args == (
vector._insert_batch,
docs[2:],
vectors[2:],
["id-3"],
2,
2,
2,
)
vector._insert_batch([], [], [], 0, 2, 1)
vector._insert_batch(docs[:1], vectors, ["id-1"], 0, 2, 1)
bad_doc = Document(page_content="doc-bad", metadata={"doc_id": "id-bad", "bad": {1}})
good_doc = Document(page_content="doc-good", metadata={"doc_id": "id-good"})
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
vector._insert_batch(
[bad_doc, good_doc],
[[0.1, 0.2], [0.3, 0.4]],
["id-bad", "id-good"],
0,
2,
1,
)
@contextmanager
def _ctx_error():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.executemany.side_effect = RuntimeError("insert failed")
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_error
with pytest.raises(RuntimeError, match="insert failed"):
vector._insert_batch([good_doc], [[0.1, 0.2]], ["id-good"], 0, 1, 1)
monkeypatch.setattr(clickzetta_module.uuid, "uuid4", lambda: "generated-id")
vector._safe_doc_id = clickzetta_module.ClickzettaVector._safe_doc_id.__get__(vector)
assert vector._safe_doc_id("") == "generated-id"
assert vector._safe_doc_id("!!!") == "generated-id"
def test_delete_by_ids_and_metadata_impl_paths(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._execute_write = MagicMock()
vector._table_exists = MagicMock(return_value=True)
vector.delete_by_ids(["id-1", "id-2"])
vector._execute_write.assert_called_once()
assert vector._execute_write.call_args.args[0] == vector._delete_by_ids_impl
vector._execute_write.reset_mock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector._execute_write.assert_called_once()
assert vector._execute_write.call_args.args[0] == vector._delete_by_metadata_field_impl
vector._safe_doc_id = MagicMock(side_effect=lambda x: x)
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
vector._delete_by_ids_impl(["id-1", "id-2"])
vector._delete_by_metadata_field_impl("document_id", "doc-1")
def test_search_by_vector_covers_cosine_and_l2_paths(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._config.vector_distance_function = "cosine_distance"
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=True)
vector._parse_metadata = MagicMock(return_value={"document_id": "doc-1", "doc_id": "seg-1"})
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.fetchall.return_value = [("seg-1", "content", '{"document_id":"doc-1"}', 0.2)]
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
cosine_docs = vector.search_by_vector(
[0.1, 0.2], top_k=3, score_threshold=0.5, document_ids_filter=["doc-1"], filter={"k": "v"}
)
assert cosine_docs[0].metadata["score"] == pytest.approx(0.9)
vector._config.vector_distance_function = "l2_distance"
l2_docs = vector.search_by_vector([0.1, 0.2], top_k=3, score_threshold=0.5)
assert l2_docs[0].metadata["score"] == pytest.approx(1 / 1.2)
def test_search_by_full_text_success_and_fallback(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=True)
@contextmanager
def _ctx_success():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.fetchall.return_value = [
("seg-1", "content-1", '"{\\"document_id\\":\\"doc-1\\"}"'),
("seg-2", "content-2", "invalid-json"),
]
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_success
docs = vector.search_by_full_text("search'value", top_k=2, document_ids_filter=["doc-1"], filter={"a": 1})
assert len(docs) == 2
assert docs[0].metadata["score"] == 1.0
assert docs[1].metadata["doc_id"] == "seg-2"
@contextmanager
def _ctx_failure():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.execute.side_effect = RuntimeError("full text failed")
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx_failure
vector._search_by_like = MagicMock(return_value=[Document(page_content="fallback", metadata={"score": 0.5})])
fallback_docs = vector.search_by_full_text("query", top_k=1)
assert fallback_docs == vector._search_by_like.return_value
def test_search_by_like_missing_table_and_delete_table(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=False)
assert vector._search_by_like("query", top_k=1) == []
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
vector.delete()
def test_clickzetta_pool_cleanup_and_shutdown_edge_paths(clickzetta_module):
pool = clickzetta_module.ClickzettaConnectionPool.__new__(clickzetta_module.ClickzettaConnectionPool)
pool._pools = {}
pool._pool_locks = {}
pool._max_pool_size = 1
pool._connection_timeout = 10
pool._lock = clickzetta_module.threading.Lock()
pool._shutdown = False
config = _config(clickzetta_module)
key = pool._get_config_key(config)
pool._pools[key] = [(MagicMock(), 1.0)]
pool._pool_locks[key] = clickzetta_module.threading.Lock()
pool._is_connection_valid = MagicMock(return_value=False)
conn = MagicMock()
pool.return_connection(config, conn)
conn.close.assert_called_once()
pool._pools["missing-lock-key"] = [(MagicMock(), 0.0)]
pool._cleanup_expired_connections()
pool.shutdown()
assert pool._shutdown is True
def test_clickzetta_pool_cleanup_thread_and_worker_exception_paths(clickzetta_module, monkeypatch):
pool = clickzetta_module.ClickzettaConnectionPool.__new__(clickzetta_module.ClickzettaConnectionPool)
pool._shutdown = False
def _cleanup_then_fail():
pool._shutdown = True
raise RuntimeError("cleanup failed")
pool._cleanup_expired_connections = MagicMock(side_effect=_cleanup_then_fail)
monkeypatch.setattr(clickzetta_module.time, "sleep", lambda _s: None)
class _Thread:
def __init__(self, target, daemon):
self._target = target
self.daemon = daemon
def start(self):
self._target()
monkeypatch.setattr(clickzetta_module.threading, "Thread", _Thread)
pool._start_cleanup_thread()
pool._cleanup_expired_connections.assert_called_once()
def test_clickzetta_parse_metadata_and_write_worker_additional_branches(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
parsed_non_dict = vector._parse_metadata("[1,2,3]", "row-1")
assert parsed_non_dict["doc_id"] == "row-1"
assert parsed_non_dict["document_id"] == "row-1"
parsed_none = vector._parse_metadata(None, "row-2")
assert parsed_none["doc_id"] == "row-2"
assert parsed_none["document_id"] == "row-2"
clickzetta_module.ClickzettaVector._shutdown = False
clickzetta_module.ClickzettaVector._write_queue = None
clickzetta_module.ClickzettaVector._write_worker()
class _BadQueue:
def get(self, timeout):
clickzetta_module.ClickzettaVector._shutdown = True
raise RuntimeError("queue failed")
clickzetta_module.ClickzettaVector._shutdown = False
clickzetta_module.ClickzettaVector._write_queue = _BadQueue()
clickzetta_module.ClickzettaVector._write_worker()
def test_clickzetta_inverted_index_existing_and_insert_non_dict_metadata(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._table_name = "table_1"
cursor = MagicMock()
cursor.fetchall.return_value = [("idx_table_1_text", "INVERTED", "page_content")]
cursor.execute.side_effect = [
None,
RuntimeError("already has index with the same type cannot create inverted index"),
None,
]
vector._create_inverted_index(cursor)
vector._safe_doc_id = MagicMock(side_effect=lambda value: str(value))
@contextmanager
def _ctx():
connection = MagicMock()
cursor_obj = MagicMock()
cursor_obj.__enter__.return_value = cursor_obj
cursor_obj.__exit__.return_value = None
connection.cursor.return_value = cursor_obj
yield connection
vector.get_connection_context = _ctx
vector._insert_batch(
[SimpleNamespace(page_content="content", metadata="not-a-dict")],
[[0.1, 0.2]],
["doc-1"],
0,
1,
1,
)
def test_clickzetta_full_text_table_missing_and_non_dict_metadata(clickzetta_module):
vector = clickzetta_module.ClickzettaVector.__new__(clickzetta_module.ClickzettaVector)
vector._config = _config(clickzetta_module)
vector._config.enable_inverted_index = True
vector._table_name = "table_1"
vector._table_exists = MagicMock(return_value=False)
assert vector.search_by_full_text("query") == []
vector._table_exists = MagicMock(return_value=True)
@contextmanager
def _ctx():
connection = MagicMock()
cursor = MagicMock()
cursor.__enter__.return_value = cursor
cursor.__exit__.return_value = None
cursor.fetchall.return_value = [
("seg-1", "content-1", "[1,2,3]"),
("seg-2", "content-2", None),
]
connection.cursor.return_value = cursor
yield connection
vector.get_connection_context = _ctx
docs = vector.search_by_full_text("query")
assert len(docs) == 2
assert docs[0].metadata["doc_id"] == "seg-1"
assert docs[1].metadata["doc_id"] == "seg-2"
@@ -1,364 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_couchbase_modules():
couchbase = types.ModuleType("couchbase")
couchbase_auth = types.ModuleType("couchbase.auth")
couchbase_cluster = types.ModuleType("couchbase.cluster")
couchbase_management = types.ModuleType("couchbase.management")
couchbase_management_search = types.ModuleType("couchbase.management.search")
couchbase_options = types.ModuleType("couchbase.options")
couchbase_vector = types.ModuleType("couchbase.vector_search")
couchbase_search = types.ModuleType("couchbase.search")
class PasswordAuthenticator:
def __init__(self, user, password):
self.user = user
self.password = password
class ClusterOptions:
def __init__(self, auth):
self.auth = auth
class SearchOptions:
def __init__(self, **kwargs):
self.kwargs = kwargs
class VectorQuery:
def __init__(self, field, vector, top_k):
self.field = field
self.vector = vector
self.top_k = top_k
class VectorSearch:
@staticmethod
def from_vector_query(vector_query):
return {"vector_query": vector_query}
class QueryStringQuery:
def __init__(self, query):
self.query = query
class SearchRequest:
@staticmethod
def create(payload):
return {"payload": payload}
class SearchIndex:
def __init__(self, name, params, source_name):
self.name = name
self.params = params
self.source_name = source_name
class _QueryResult:
def __init__(self, rows=None):
self._rows = rows or []
def execute(self):
return self
def __iter__(self):
return iter(self._rows)
class _SearchIter:
def __init__(self, rows=None):
self._rows = rows or []
def rows(self):
return self._rows
class _Collection:
def __init__(self):
self.upsert = MagicMock(return_value=True)
class _SearchIndexManager:
def __init__(self):
self.upsert_index = MagicMock()
class _Scope:
def __init__(self):
self._collection = _Collection()
self._search_index_manager = _SearchIndexManager()
self.search = MagicMock(return_value=_SearchIter())
def collection(self, _name):
return self._collection
def search_indexes(self):
return self._search_index_manager
class _CollectionManager:
def __init__(self):
self.create_collection = MagicMock()
self.drop_collection = MagicMock()
self.get_all_scopes = MagicMock(return_value=[])
class _Bucket:
def __init__(self):
self._scope = _Scope()
self._collections = _CollectionManager()
def scope(self, _scope_name):
return self._scope
def collections(self):
return self._collections
class Cluster:
def __init__(self, connection_string, options):
self.connection_string = connection_string
self.options = options
self._bucket = _Bucket()
self.wait_until_ready = MagicMock()
self.query = MagicMock(return_value=_QueryResult())
def bucket(self, _name):
return self._bucket
couchbase_auth.PasswordAuthenticator = PasswordAuthenticator
couchbase_cluster.Cluster = Cluster
couchbase_management_search.SearchIndex = SearchIndex
couchbase_options.ClusterOptions = ClusterOptions
couchbase_options.SearchOptions = SearchOptions
couchbase_vector.VectorQuery = VectorQuery
couchbase_vector.VectorSearch = VectorSearch
couchbase_search.QueryStringQuery = QueryStringQuery
couchbase_search.SearchRequest = SearchRequest
couchbase.search = couchbase_search
couchbase.management = couchbase_management
return {
"couchbase": couchbase,
"couchbase.auth": couchbase_auth,
"couchbase.cluster": couchbase_cluster,
"couchbase.management": couchbase_management,
"couchbase.management.search": couchbase_management_search,
"couchbase.options": couchbase_options,
"couchbase.vector_search": couchbase_vector,
"couchbase.search": couchbase_search,
}
@pytest.fixture
def couchbase_module(monkeypatch):
for name, module in _build_fake_couchbase_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.couchbase.couchbase_vector as module
return importlib.reload(module)
def _config(module):
return module.CouchbaseConfig(
connection_string="couchbase://localhost",
user="user",
password="pass",
bucket_name="bucket",
scope_name="scope",
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("connection_string", "", "CONNECTION_STRING is required"),
("user", "", "COUCHBASE_USER is required"),
("password", "", "COUCHBASE_PASSWORD is required"),
("bucket_name", "", "COUCHBASE_PASSWORD is required"),
("scope_name", "", "COUCHBASE_SCOPE_NAME is required"),
],
)
def test_couchbase_config_validation(couchbase_module, field, value, message):
values = _config(couchbase_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
couchbase_module.CouchbaseConfig.model_validate(values)
def test_init_sets_cluster_handles(couchbase_module):
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
assert vector._bucket_name == "bucket"
assert vector._scope_name == "scope"
vector._cluster.wait_until_ready.assert_called_once()
def test_create_and_create_collection_branches(couchbase_module, monkeypatch):
vector = couchbase_module.CouchbaseVector.__new__(couchbase_module.CouchbaseVector)
vector._collection_name = "collection_1"
vector._client_config = _config(couchbase_module)
vector._scope_name = "scope"
vector._bucket_name = "bucket"
vector._bucket = MagicMock()
vector._scope = MagicMock()
vector._collection_exists = MagicMock(return_value=False)
vector.add_texts = MagicMock()
monkeypatch.setattr(couchbase_module.uuid, "uuid4", lambda: "a-b-c")
vector._create_collection = MagicMock()
docs = [Document(page_content="text", metadata={"doc_id": "id-1"})]
vector.create(docs, [[0.1, 0.2]])
vector._create_collection.assert_called_once_with(uuid="abc", vector_length=2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(couchbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(couchbase_module.redis_client, "set", MagicMock())
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
monkeypatch.setattr(couchbase_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(vector_length=2, uuid="uuid-1")
vector._bucket.collections().create_collection.assert_not_called()
monkeypatch.setattr(couchbase_module.redis_client, "get", MagicMock(return_value=None))
vector._collection_exists = MagicMock(return_value=True)
vector._create_collection(vector_length=2, uuid="uuid-2")
vector._bucket.collections().create_collection.assert_not_called()
vector._collection_exists = MagicMock(return_value=False)
vector._create_collection(vector_length=3, uuid="uuid-3")
vector._bucket.collections().create_collection.assert_called_once_with("scope", "collection_1")
vector._scope.search_indexes().upsert_index.assert_called_once()
search_index = vector._scope.search_indexes().upsert_index.call_args.args[0]
assert search_index.name == "collection_1_search"
assert (
search_index.params["mapping"]["types"]["scope.collection_1"]["properties"]["embedding"]["fields"][0]["dims"]
== 3
)
couchbase_module.redis_client.set.assert_called_once()
def test_collection_exists_get_type_and_add_texts(couchbase_module):
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
scope_obj = SimpleNamespace(name="scope", collections=[SimpleNamespace(name="collection_1")])
vector._bucket.collections().get_all_scopes.return_value = [scope_obj]
assert vector._collection_exists("collection_1") is True
scope_obj = SimpleNamespace(name="scope", collections=[SimpleNamespace(name="other")])
vector._bucket.collections().get_all_scopes.return_value = [scope_obj]
assert vector._collection_exists("collection_1") is False
vector._get_uuids = MagicMock(return_value=["id-1", "id-2"])
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "id-2"]
assert vector._scope.collection("collection_1").upsert.call_count == 2
assert vector.get_type() == couchbase_module.VectorType.COUCHBASE
def test_query_delete_helpers(couchbase_module):
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
vector._cluster.query.return_value = SimpleNamespace(execute=lambda: iter([{"count": 2}]))
assert vector.text_exists("id-1") is True
vector._cluster.query.return_value = SimpleNamespace(execute=lambda: iter([]))
assert vector.text_exists("id-2") is False
query_result = MagicMock()
query_result.execute.return_value = None
vector._cluster.query.return_value = query_result
vector.delete_by_ids(["id-1", "id-2"])
vector.delete_by_document_id("id-1")
vector.delete_by_metadata_field("document_id", "doc-1")
assert vector._cluster.query.call_count >= 3
vector._cluster.query.side_effect = RuntimeError("delete failed")
vector.delete_by_ids(["id-3"])
def test_search_methods_and_format_metadata(couchbase_module):
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
row_1 = SimpleNamespace(fields={"text": "doc-a", "metadata.document_id": "d-1"}, score=0.9)
row_2 = SimpleNamespace(fields={"text": "doc-b", "metadata.document_id": "d-2"}, score=0.3)
vector._scope.search.return_value = SimpleNamespace(rows=lambda: [row_1, row_2])
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].page_content == "doc-a"
assert docs[0].metadata["document_id"] == "d-1"
assert docs[0].metadata["score"] == pytest.approx(0.9)
vector._scope.search.side_effect = RuntimeError("search error")
with pytest.raises(ValueError, match="Search failed"):
vector.search_by_vector([0.1], top_k=1)
vector._scope.search.side_effect = None
row_3 = SimpleNamespace(fields={"text": "full-text", "metadata.doc_id": "x"}, score=0.7)
vector._scope.search.return_value = SimpleNamespace(rows=lambda: [row_3])
docs = vector.search_by_full_text("hello", top_k=1)
assert len(docs) == 1
assert docs[0].metadata["doc_id"] == "x"
vector._scope.search.side_effect = RuntimeError("full text failed")
with pytest.raises(ValueError, match="Search failed"):
vector.search_by_full_text("hello", top_k=1)
assert vector._format_metadata({"metadata.a": 1, "plain": 2}) == {"a": 1, "plain": 2}
def test_delete_collection_and_factory(couchbase_module, monkeypatch):
vector = couchbase_module.CouchbaseVector("collection_1", _config(couchbase_module))
scopes = [
SimpleNamespace(collections=[SimpleNamespace(name="other")]),
SimpleNamespace(collections=[SimpleNamespace(name="collection_1")]),
]
vector._bucket.collections().get_all_scopes.return_value = scopes
vector.delete()
vector._bucket.collections().drop_collection.assert_called_once_with("_default", "collection_1")
factory = couchbase_module.CouchbaseVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(couchbase_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(
couchbase_module,
"current_app",
SimpleNamespace(
config={
"COUCHBASE_CONNECTION_STRING": "couchbase://localhost",
"COUCHBASE_USER": "user",
"COUCHBASE_PASSWORD": "pass",
"COUCHBASE_BUCKET_NAME": "bucket",
"COUCHBASE_SCOPE_NAME": "scope",
}
),
)
with patch.object(couchbase_module, "CouchbaseVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,121 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
def _build_fake_elasticsearch_modules():
elasticsearch = types.ModuleType("elasticsearch")
class ConnectionError(Exception):
pass
class Elasticsearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.ping = MagicMock(return_value=True)
self.info = MagicMock(return_value={"version": {"number": "8.12.0"}})
self.indices = SimpleNamespace(
refresh=MagicMock(), delete=MagicMock(), exists=MagicMock(return_value=False), create=MagicMock()
)
elasticsearch.Elasticsearch = Elasticsearch
elasticsearch.ConnectionError = ConnectionError
return {"elasticsearch": elasticsearch}
@pytest.fixture
def elasticsearch_ja_module(monkeypatch):
for name, module in _build_fake_elasticsearch_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.elasticsearch.elasticsearch_ja_vector as ja_module
import core.rag.datasource.vdb.elasticsearch.elasticsearch_vector as base_module
importlib.reload(base_module)
return importlib.reload(ja_module)
def test_create_collection_cache_hit(elasticsearch_ja_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "set", MagicMock())
vector = elasticsearch_ja_module.ElasticSearchJaVector.__new__(elasticsearch_ja_module.ElasticSearchJaVector)
vector._collection_name = "test"
vector._client = MagicMock()
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_not_called()
elasticsearch_ja_module.redis_client.set.assert_not_called()
def test_create_collection_create_and_exists_paths(elasticsearch_ja_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(elasticsearch_ja_module.redis_client, "set", MagicMock())
vector = elasticsearch_ja_module.ElasticSearchJaVector.__new__(elasticsearch_ja_module.ElasticSearchJaVector)
vector._collection_name = "test"
vector._client = MagicMock()
vector._client.indices.exists.return_value = False
vector.create_collection([[0.1, 0.2, 0.3]], [{}])
vector._client.indices.create.assert_called_once()
kwargs = vector._client.indices.create.call_args.kwargs
assert kwargs["index"] == "test"
assert kwargs["mappings"]["properties"][elasticsearch_ja_module.Field.VECTOR]["dims"] == 3
elasticsearch_ja_module.redis_client.set.assert_called_once()
vector._client.indices.create.reset_mock()
elasticsearch_ja_module.redis_client.set.reset_mock()
vector._client.indices.exists.return_value = True
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_not_called()
elasticsearch_ja_module.redis_client.set.assert_called_once()
def test_ja_factory_uses_existing_or_generated_collection(elasticsearch_ja_module, monkeypatch):
factory = elasticsearch_ja_module.ElasticSearchJaVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(elasticsearch_ja_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(
elasticsearch_ja_module,
"current_app",
SimpleNamespace(
config={
"ELASTICSEARCH_HOST": "localhost",
"ELASTICSEARCH_PORT": 9200,
"ELASTICSEARCH_USERNAME": "elastic",
"ELASTICSEARCH_PASSWORD": "secret",
}
),
)
with patch.object(elasticsearch_ja_module, "ElasticSearchJaVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["index_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["index_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,405 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_elasticsearch_modules():
elasticsearch = types.ModuleType("elasticsearch")
class ConnectionError(Exception):
pass
class Elasticsearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.ping = MagicMock(return_value=True)
self.info = MagicMock(return_value={"version": {"number": "8.12.0-SNAPSHOT"}})
self.index = MagicMock()
self.exists = MagicMock(return_value=False)
self.delete = MagicMock()
self.search = MagicMock(return_value={"hits": {"hits": []}})
self.indices = SimpleNamespace(
refresh=MagicMock(),
delete=MagicMock(),
exists=MagicMock(return_value=False),
create=MagicMock(),
)
elasticsearch.Elasticsearch = Elasticsearch
elasticsearch.ConnectionError = ConnectionError
return {"elasticsearch": elasticsearch}
@pytest.fixture
def elasticsearch_module(monkeypatch):
for name, module in _build_fake_elasticsearch_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.elasticsearch.elasticsearch_vector as module
return importlib.reload(module)
def _regular_config(module, **overrides):
values = {
"host": "localhost",
"port": 9200,
"username": "elastic",
"password": "secret",
"verify_certs": False,
"request_timeout": 10,
"retry_on_timeout": True,
"max_retries": 3,
}
values.update(overrides)
return module.ElasticSearchConfig.model_validate(values)
def _cloud_config(module, **overrides):
values = {
"use_cloud": True,
"cloud_url": "https://cloud.example:9243",
"api_key": "api-key",
"verify_certs": True,
"ca_certs": "/tmp/ca.pem",
"request_timeout": 10,
"retry_on_timeout": True,
"max_retries": 3,
}
values.update(overrides)
return module.ElasticSearchConfig.model_validate(values)
@pytest.mark.parametrize(
("values", "message"),
[
({"use_cloud": True, "cloud_url": None, "api_key": "x"}, "cloud_url is required"),
({"use_cloud": True, "cloud_url": "https://cloud", "api_key": None}, "api_key is required"),
({"host": None, "port": 9200, "username": "u", "password": "p"}, "HOST is required"),
({"host": "h", "port": None, "username": "u", "password": "p"}, "PORT is required"),
({"host": "h", "port": 9200, "username": None, "password": "p"}, "USERNAME is required"),
({"host": "h", "port": 9200, "username": "u", "password": None}, "PASSWORD is required"),
],
)
def test_elasticsearch_config_validation(elasticsearch_module, values, message):
with pytest.raises(ValidationError, match=message):
elasticsearch_module.ElasticSearchConfig.model_validate(values)
def test_init_client_cloud_configuration(elasticsearch_module):
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
client = MagicMock()
client.ping.return_value = True
with patch.object(elasticsearch_module, "Elasticsearch", return_value=client) as es_cls:
result = vector._init_client(_cloud_config(elasticsearch_module))
assert result is client
kwargs = es_cls.call_args.kwargs
assert kwargs["hosts"] == ["https://cloud.example:9243"]
assert kwargs["api_key"] == "api-key"
assert kwargs["verify_certs"] is True
assert kwargs["ca_certs"] == "/tmp/ca.pem"
def test_init_client_regular_https_and_http_fallback(elasticsearch_module):
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
client = MagicMock()
client.ping.return_value = True
with patch.object(elasticsearch_module, "Elasticsearch", return_value=client) as es_cls:
vector._init_client(
_regular_config(
elasticsearch_module,
host="https://es.example",
port=9443,
verify_certs=True,
ca_certs="/tmp/ca.pem",
)
)
kwargs = es_cls.call_args.kwargs
assert kwargs["hosts"] == ["https://es.example:9443"]
assert kwargs["verify_certs"] is True
assert kwargs["ca_certs"] == "/tmp/ca.pem"
with patch.object(elasticsearch_module, "Elasticsearch", return_value=client) as es_cls:
vector._init_client(_regular_config(elasticsearch_module, host="es.internal", port=9200))
kwargs = es_cls.call_args.kwargs
assert kwargs["hosts"] == ["http://es.internal:9200"]
assert "verify_certs" not in kwargs
def test_init_client_connection_failures(elasticsearch_module):
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
client = MagicMock()
client.ping.return_value = False
with patch.object(elasticsearch_module, "Elasticsearch", return_value=client):
with pytest.raises(ConnectionError, match="Failed to connect"):
vector._init_client(_regular_config(elasticsearch_module))
with patch.object(
elasticsearch_module,
"Elasticsearch",
side_effect=elasticsearch_module.ElasticsearchConnectionError("boom"),
):
with pytest.raises(ConnectionError, match="Vector database connection error"):
vector._init_client(_regular_config(elasticsearch_module))
with patch.object(elasticsearch_module, "Elasticsearch", side_effect=RuntimeError("oops")):
with pytest.raises(ConnectionError, match="initialization failed"):
vector._init_client(_regular_config(elasticsearch_module))
def test_init_get_version_and_check_version(elasticsearch_module):
with (
patch.object(elasticsearch_module.ElasticSearchVector, "_init_client", return_value=MagicMock()) as init_client,
patch.object(elasticsearch_module.ElasticSearchVector, "_get_version", return_value="8.10.0") as get_version,
patch.object(elasticsearch_module.ElasticSearchVector, "_check_version") as check_version,
):
vector = elasticsearch_module.ElasticSearchVector(
"collection_1", _regular_config(elasticsearch_module), attributes=["doc_id"]
)
init_client.assert_called_once()
get_version.assert_called_once()
check_version.assert_called_once()
assert vector._attributes == ["doc_id"]
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
vector._client = MagicMock()
vector._client.info.return_value = {"version": {"number": "8.13.2-SNAPSHOT"}}
assert vector._get_version() == "8.13.2"
vector._version = "7.17.0"
with pytest.raises(ValueError, match="greater than 8.0.0"):
vector._check_version()
vector._version = "8.0.0"
vector._check_version()
def test_crud_methods_and_get_type(elasticsearch_module):
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.indices = SimpleNamespace(refresh=MagicMock(), delete=MagicMock())
vector._get_uuids = MagicMock(return_value=["id-1", "id-2"])
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "id-2"]
assert vector._client.index.call_count == 2
vector._client.indices.refresh.assert_called_once_with(index="collection_1")
vector._client.exists.return_value = True
assert vector.text_exists("id-1") is True
vector.delete_by_ids([])
vector._client.delete.assert_not_called()
vector.delete_by_ids(["id-1", "id-2"])
assert vector._client.delete.call_count == 2
vector._client.search.return_value = {"hits": {"hits": [{"_id": "id-1"}]}}
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("doc_id", "d1")
vector.delete_by_ids.assert_called_once_with(["id-1"])
vector.delete_by_ids.reset_mock()
vector._client.search.return_value = {"hits": {"hits": []}}
vector.delete_by_metadata_field("doc_id", "d2")
vector.delete_by_ids.assert_not_called()
vector.delete()
vector._client.indices.delete.assert_called_once_with(index="collection_1")
assert vector.get_type() == elasticsearch_module.VectorType.ELASTICSEARCH
def test_search_by_vector_and_full_text(elasticsearch_module):
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_score": 0.8,
"_source": {
elasticsearch_module.Field.CONTENT_KEY: "doc-a",
elasticsearch_module.Field.VECTOR: [0.1],
elasticsearch_module.Field.METADATA_KEY: {"doc_id": "1", "document_id": "d-1"},
},
},
{
"_score": 0.2,
"_source": {
elasticsearch_module.Field.CONTENT_KEY: "doc-b",
elasticsearch_module.Field.VECTOR: [0.2],
elasticsearch_module.Field.METADATA_KEY: {"doc_id": "2", "document_id": "d-2"},
},
},
]
}
}
docs = vector.search_by_vector(
[0.1, 0.2],
top_k=2,
score_threshold=0.5,
document_ids_filter=["d-1", "d-2"],
)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.8)
knn = vector._client.search.call_args.kwargs["knn"]
assert knn["k"] == 2
assert knn["num_candidates"] == 3
assert "filter" in knn
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_source": {
elasticsearch_module.Field.CONTENT_KEY: "text-hit",
elasticsearch_module.Field.VECTOR: [0.3],
elasticsearch_module.Field.METADATA_KEY: {"doc_id": "3"},
}
}
]
}
}
docs = vector.search_by_full_text("hello", top_k=3, document_ids_filter=["d-3"])
assert len(docs) == 1
assert docs[0].page_content == "text-hit"
query = vector._client.search.call_args.kwargs["query"]
assert "bool" in query
def test_create_and_create_collection_paths(elasticsearch_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(elasticsearch_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(elasticsearch_module.redis_client, "set", MagicMock())
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.indices = SimpleNamespace(exists=MagicMock(return_value=False), create=MagicMock())
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="a", metadata={"doc_id": "1"})]
vector.create(docs, [[0.1]])
vector.create_collection.assert_called_once()
vector.add_texts.assert_called_once_with(docs, [[0.1]])
vector = elasticsearch_module.ElasticSearchVector.__new__(elasticsearch_module.ElasticSearchVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.indices = SimpleNamespace(exists=MagicMock(return_value=False), create=MagicMock())
monkeypatch.setattr(elasticsearch_module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_not_called()
monkeypatch.setattr(elasticsearch_module.redis_client, "get", MagicMock(return_value=None))
vector._client.indices.exists.return_value = False
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_called_once()
mappings = vector._client.indices.create.call_args.kwargs["mappings"]
assert mappings["properties"][elasticsearch_module.Field.VECTOR]["dims"] == 2
elasticsearch_module.redis_client.set.assert_called_once()
vector._client.indices.create.reset_mock()
elasticsearch_module.redis_client.set.reset_mock()
vector._client.indices.exists.return_value = True
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_not_called()
elasticsearch_module.redis_client.set.assert_called_once()
def test_elasticsearch_factory_branches(elasticsearch_module, monkeypatch):
factory = elasticsearch_module.ElasticSearchVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(elasticsearch_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(
elasticsearch_module,
"current_app",
SimpleNamespace(
config={
"ELASTICSEARCH_USE_CLOUD": False,
"ELASTICSEARCH_HOST": "es-host",
"ELASTICSEARCH_PORT": 9200,
"ELASTICSEARCH_USERNAME": "elastic",
"ELASTICSEARCH_PASSWORD": "secret",
"ELASTICSEARCH_VERIFY_CERTS": False,
}
),
)
with patch.object(elasticsearch_module, "ElasticSearchVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
cfg = vector_cls.call_args.kwargs["config"]
assert cfg.use_cloud is False
assert vector_cls.call_args.kwargs["index_name"] == "EXISTING_COLLECTION"
monkeypatch.setattr(
elasticsearch_module,
"current_app",
SimpleNamespace(
config={
"ELASTICSEARCH_USE_CLOUD": True,
"ELASTICSEARCH_CLOUD_URL": "https://cloud.elastic",
"ELASTICSEARCH_API_KEY": "api-key",
"ELASTICSEARCH_VERIFY_CERTS": True,
}
),
)
with patch.object(elasticsearch_module, "ElasticSearchVector", return_value="vector") as vector_cls:
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_2 == "vector"
cfg = vector_cls.call_args.kwargs["config"]
assert cfg.use_cloud is True
assert cfg.cloud_url == "https://cloud.elastic"
assert dataset_without_index.index_struct is not None
monkeypatch.setattr(
elasticsearch_module,
"current_app",
SimpleNamespace(
config={
"ELASTICSEARCH_USE_CLOUD": True,
"ELASTICSEARCH_CLOUD_URL": None,
"ELASTICSEARCH_HOST": "fallback-host",
"ELASTICSEARCH_PORT": 9201,
"ELASTICSEARCH_USERNAME": "elastic",
"ELASTICSEARCH_PASSWORD": "secret",
}
),
)
with patch.object(elasticsearch_module, "ElasticSearchVector", return_value="vector") as vector_cls:
factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
cfg = vector_cls.call_args.kwargs["config"]
assert cfg.use_cloud is False
assert cfg.host == "fallback-host"
@@ -1,371 +0,0 @@
import importlib
import json
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_hologres_modules():
holo_module = types.ModuleType("holo_search_sdk")
holo_types_module = types.ModuleType("holo_search_sdk.types")
holo_types_module.BaseQuantizationType = str
holo_types_module.DistanceType = str
holo_types_module.TokenizerType = str
def _connect(**kwargs):
client = MagicMock()
client.kwargs = kwargs
client.connect = MagicMock()
client.check_table_exist = MagicMock(return_value=False)
client.open_table = MagicMock(return_value=MagicMock())
client.execute = MagicMock(return_value=[])
client.drop_table = MagicMock()
return client
holo_module.connect = MagicMock(side_effect=_connect)
return {
"holo_search_sdk": holo_module,
"holo_search_sdk.types": holo_types_module,
}
@pytest.fixture
def hologres_module(monkeypatch):
for name, module in _build_fake_hologres_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.hologres.hologres_vector as module
return importlib.reload(module)
def _valid_config(module):
return module.HologresVectorConfig(
host="localhost",
port=80,
database="dify",
access_key_id="ak",
access_key_secret="sk",
schema_name="public",
tokenizer="jieba",
distance_method="Cosine",
base_quantization_type="rabitq",
max_degree=64,
ef_construction=400,
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config HOLOGRES_HOST is required"),
("database", "", "config HOLOGRES_DATABASE is required"),
("access_key_id", "", "config HOLOGRES_ACCESS_KEY_ID is required"),
("access_key_secret", "", "config HOLOGRES_ACCESS_KEY_SECRET is required"),
],
)
def test_hologres_config_validation(hologres_module, field, value, message):
values = _valid_config(hologres_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
hologres_module.HologresVectorConfig.model_validate(values)
def test_init_client_and_get_type(hologres_module):
vector = hologres_module.HologresVector("Collection_One", _valid_config(hologres_module))
hologres_module.holo.connect.assert_called_once_with(
host="localhost",
port=80,
database="dify",
access_key_id="ak",
access_key_secret="sk",
schema="public",
)
vector._client.connect.assert_called_once()
assert vector.table_name == "embedding_collection_one"
assert vector.get_type() == hologres_module.VectorType.HOLOGRES
def test_create_delegates_collection_creation_and_upsert(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="hello", metadata={"doc_id": "seg-1"})]
result = vector.create(docs, [[0.1, 0.2]])
assert result is None
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_returns_empty_for_empty_documents(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
assert vector.add_texts([], []) == []
vector._client.open_table.assert_not_called()
def test_add_texts_batches_and_serializes_metadata(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
table = vector._client.open_table.return_value
documents = [
Document(page_content=f"doc-{i}", metadata={"doc_id": f"id-{i}", "document_id": f"document-{i}"})
for i in range(100)
]
documents.append(SimpleNamespace(page_content="doc-100", metadata=None))
embeddings = [[float(i)] for i in range(len(documents))]
ids = vector.add_texts(documents, embeddings)
assert ids[:2] == ["id-0", "id-1"]
assert ids[-1] == ""
assert len(ids) == 101
assert vector._client.open_table.call_count == 2
assert table.upsert_multi.call_count == 2
first_call = table.upsert_multi.call_args_list[0].kwargs
second_call = table.upsert_multi.call_args_list[1].kwargs
assert first_call["index_column"] == "id"
assert first_call["column_names"] == ["id", "text", "meta", "embedding"]
assert first_call["update_columns"] == ["text", "meta", "embedding"]
assert len(first_call["values"]) == 100
assert json.loads(first_call["values"][0][2]) == {"doc_id": "id-0", "document_id": "document-0"}
assert second_call["values"][0][0] == ""
assert second_call["values"][0][2] == "{}"
def test_text_exists_handles_missing_and_present_tables(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.side_effect = [False, True]
vector._client.execute.return_value = [(1,)]
assert vector.text_exists("seg-1") is False
assert vector.text_exists("seg-1") is True
vector._client.execute.assert_called_once()
def test_get_ids_by_metadata_field_returns_ids_or_none(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.execute.side_effect = [[("id-1",), ("id-2",)], []]
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"]
assert vector.get_ids_by_metadata_field("document_id", "doc-1") is None
def test_delete_by_ids_branches(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector.delete_by_ids([])
vector._client.check_table_exist.assert_not_called()
vector._client.check_table_exist.return_value = False
vector.delete_by_ids(["id-1"])
vector._client.execute.assert_not_called()
vector._client.check_table_exist.return_value = True
vector.delete_by_ids(["id-1", "id-2"])
vector._client.execute.assert_called_once()
def test_delete_by_metadata_field_branches(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.return_value = False
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.execute.assert_not_called()
vector._client.check_table_exist.return_value = True
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.execute.assert_called_once()
def test_search_by_vector_returns_empty_when_table_missing(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.return_value = False
assert vector.search_by_vector([0.1, 0.2]) == []
def test_search_by_vector_applies_filter_and_processes_results(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.return_value = True
table = vector._client.open_table.return_value
query = MagicMock()
table.search_vector.return_value = query
query.select.return_value = query
query.limit.return_value = query
query.where.return_value = query
query.fetchall.return_value = [
(0.2, "seg-1", "doc-1", '{"doc_id":"seg-1","document_id":"doc-1"}'),
(0.9, "seg-2", "doc-2", {"doc_id": "seg-2", "document_id": "doc-2"}),
]
docs = vector.search_by_vector(
[0.1, 0.2],
top_k=2,
score_threshold=0.5,
document_ids_filter=["doc-1"],
)
assert len(docs) == 1
assert docs[0].page_content == "doc-1"
assert docs[0].metadata["doc_id"] == "seg-1"
assert docs[0].metadata["score"] == pytest.approx(0.8)
table.search_vector.assert_called_once()
query.where.assert_called_once()
def test_search_by_full_text_returns_empty_when_table_missing(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.return_value = False
assert vector.search_by_full_text("query") == []
def test_search_by_full_text_applies_filter_and_processes_results(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.return_value = True
table = vector._client.open_table.return_value
search_query = MagicMock()
table.search_text.return_value = search_query
search_query.limit.return_value = search_query
search_query.where.return_value = search_query
search_query.fetchall.return_value = [
("seg-1", "doc-1", '{"doc_id":"seg-1"}', [0.1], 0.95),
("seg-2", "doc-2", {"doc_id": "seg-2"}, [0.2], 0.7),
]
docs = vector.search_by_full_text("query", top_k=2, document_ids_filter=["doc-1"])
assert len(docs) == 2
assert docs[0].metadata["doc_id"] == "seg-1"
assert docs[0].metadata["score"] == pytest.approx(0.95)
assert docs[1].metadata["score"] == pytest.approx(0.7)
table.search_text.assert_called_once()
search_query.where.assert_called_once()
def test_delete_handles_existing_and_missing_tables(hologres_module):
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.side_effect = [False, True]
vector.delete()
vector._client.drop_table.assert_not_called()
vector.delete()
vector._client.drop_table.assert_called_once_with(vector.table_name)
def test_create_collection_returns_early_when_cache_hits(hologres_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = False
monkeypatch.setattr(hologres_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(hologres_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(hologres_module.redis_client, "set", MagicMock())
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._create_collection(3)
vector._client.check_table_exist.assert_not_called()
hologres_module.redis_client.set.assert_not_called()
def test_create_collection_creates_table_and_indexes(hologres_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = False
monkeypatch.setattr(hologres_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(hologres_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(hologres_module.redis_client, "set", MagicMock())
monkeypatch.setattr(hologres_module.time, "sleep", MagicMock())
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.side_effect = [False, False, True]
table = vector._client.open_table.return_value
vector._create_collection(3)
vector._client.execute.assert_called_once()
table.set_vector_index.assert_called_once_with(
column="embedding",
distance_method="Cosine",
base_quantization_type="rabitq",
max_degree=64,
ef_construction=400,
use_reorder=True,
)
table.create_text_index.assert_called_once_with(
index_name="ft_idx_collection_one",
column="text",
tokenizer="jieba",
)
hologres_module.redis_client.set.assert_called_once()
def test_create_collection_raises_when_table_never_becomes_ready(hologres_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = False
monkeypatch.setattr(hologres_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(hologres_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(hologres_module.redis_client, "set", MagicMock())
monkeypatch.setattr(hologres_module.time, "sleep", MagicMock())
vector = hologres_module.HologresVector("collection_one", _valid_config(hologres_module))
vector._client.check_table_exist.side_effect = [False] + [False] * 15
with pytest.raises(RuntimeError, match="was not ready after 30s"):
vector._create_collection(3)
hologres_module.redis_client.set.assert_not_called()
def test_hologres_factory_uses_existing_or_generated_collection(hologres_module, monkeypatch):
factory = hologres_module.HologresVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "existing_collection"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(hologres_module.Dataset, "gen_collection_name_by_id", lambda _id: "generated_collection")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_HOST", "127.0.0.1")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_PORT", 80)
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_DATABASE", "dify")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_ACCESS_KEY_ID", "ak")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_ACCESS_KEY_SECRET", "sk")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_SCHEMA", "public")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_TOKENIZER", "jieba")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_DISTANCE_METHOD", "Cosine")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_BASE_QUANTIZATION_TYPE", "rabitq")
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_MAX_DEGREE", 64)
monkeypatch.setattr(hologres_module.dify_config, "HOLOGRES_EF_CONSTRUCTION", 400)
with patch.object(hologres_module, "HologresVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "generated_collection"
generated_config = vector_cls.call_args_list[1].kwargs["config"]
assert generated_config.host == "127.0.0.1"
assert generated_config.database == "dify"
assert generated_config.access_key_id == "ak"
assert json.loads(dataset_without_index.index_struct) == {
"type": hologres_module.VectorType.HOLOGRES,
"vector_store": {"class_prefix": "generated_collection"},
}
@@ -1,243 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_elasticsearch_modules():
elasticsearch = types.ModuleType("elasticsearch")
class Elasticsearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.index = MagicMock()
self.exists = MagicMock(return_value=False)
self.delete = MagicMock()
self.search = MagicMock(return_value={"hits": {"hits": []}})
self.indices = SimpleNamespace(
refresh=MagicMock(), delete=MagicMock(), exists=MagicMock(return_value=False), create=MagicMock()
)
elasticsearch.Elasticsearch = Elasticsearch
return {"elasticsearch": elasticsearch}
@pytest.fixture
def huawei_module(monkeypatch):
for name, module in _build_fake_elasticsearch_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.huawei.huawei_cloud_vector as module
return importlib.reload(module)
def _config(module):
return module.HuaweiCloudVectorConfig(hosts="http://localhost:9200", username="user", password="pass")
def test_create_ssl_context(huawei_module):
ctx = huawei_module.create_ssl_context()
assert ctx.check_hostname is False
assert ctx.verify_mode == huawei_module.ssl.CERT_NONE
def test_huawei_config_validation_and_params(huawei_module):
with pytest.raises(ValidationError, match="HOSTS is required"):
huawei_module.HuaweiCloudVectorConfig.model_validate({"hosts": ""})
config = _config(huawei_module)
params = config.to_elasticsearch_params()
assert params["hosts"] == ["http://localhost:9200"]
assert params["basic_auth"] == ("user", "pass")
config = huawei_module.HuaweiCloudVectorConfig(hosts="host1,host2", username=None, password=None)
params = config.to_elasticsearch_params()
assert "basic_auth" not in params
def test_init_get_type_and_add_texts(huawei_module):
vector = huawei_module.HuaweiCloudVector("COLLECTION", _config(huawei_module))
assert vector._collection_name == "collection"
assert vector.get_type() == huawei_module.VectorType.HUAWEI_CLOUD
vector._get_uuids = MagicMock(return_value=["id-1", "id-2"])
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "id-2"]
assert vector._client.index.call_count == 2
vector._client.indices.refresh.assert_called_once_with(index="collection")
def test_crud_methods(huawei_module):
vector = huawei_module.HuaweiCloudVector("collection", _config(huawei_module))
vector._client.exists.return_value = True
assert vector.text_exists("id-1") is True
vector.delete_by_ids([])
vector._client.delete.assert_not_called()
vector.delete_by_ids(["id-1"])
vector._client.delete.assert_called_once_with(index="collection", id="id-1")
vector._client.search.return_value = {"hits": {"hits": [{"_id": "id-1"}]}}
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("doc_id", "x")
vector.delete_by_ids.assert_called_once_with(["id-1"])
vector.delete_by_ids.reset_mock()
vector._client.search.return_value = {"hits": {"hits": []}}
vector.delete_by_metadata_field("doc_id", "x")
vector.delete_by_ids.assert_not_called()
vector.delete()
vector._client.indices.delete.assert_called_once_with(index="collection")
def test_search_by_vector_and_full_text(huawei_module):
vector = huawei_module.HuaweiCloudVector("collection", _config(huawei_module))
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_score": 0.9,
"_source": {
huawei_module.Field.CONTENT_KEY: "doc-a",
huawei_module.Field.VECTOR: [0.1],
huawei_module.Field.METADATA_KEY: {"doc_id": "1"},
},
},
{
"_score": 0.1,
"_source": {
huawei_module.Field.CONTENT_KEY: "doc-b",
huawei_module.Field.VECTOR: [0.2],
huawei_module.Field.METADATA_KEY: {"doc_id": "2"},
},
},
]
}
}
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
query_body = vector._client.search.call_args.kwargs["body"]
assert query_body["query"]["vector"][huawei_module.Field.VECTOR]["topk"] == 2
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_source": {
huawei_module.Field.CONTENT_KEY: "text-hit",
huawei_module.Field.VECTOR: [0.3],
huawei_module.Field.METADATA_KEY: {"doc_id": "3"},
}
}
]
}
}
docs = vector.search_by_full_text("hello", top_k=3)
assert len(docs) == 1
assert docs[0].page_content == "text-hit"
def test_search_by_vector_skips_hits_without_metadata(huawei_module, monkeypatch):
class FakeDocument:
def __init__(self, page_content, vector, metadata):
self.page_content = page_content
self.vector = vector
self.metadata = None
monkeypatch.setattr(huawei_module, "Document", FakeDocument)
vector = huawei_module.HuaweiCloudVector("collection", _config(huawei_module))
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_score": 0.9,
"_source": {
huawei_module.Field.CONTENT_KEY: "doc-a",
huawei_module.Field.VECTOR: [0.1],
huawei_module.Field.METADATA_KEY: {"doc_id": "1"},
},
}
]
}
}
docs = vector.search_by_vector([0.1, 0.2], top_k=1, score_threshold=0.5)
assert docs == []
def test_create_and_create_collection_paths(huawei_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(huawei_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(huawei_module.redis_client, "set", MagicMock())
vector = huawei_module.HuaweiCloudVector("collection", _config(huawei_module))
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="a", metadata={"doc_id": "1"})]
vector.create(docs, [[0.1]])
vector.create_collection.assert_called_once()
vector.add_texts.assert_called_once_with(docs, [[0.1]])
vector = huawei_module.HuaweiCloudVector("collection", _config(huawei_module))
monkeypatch.setattr(huawei_module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_not_called()
monkeypatch.setattr(huawei_module.redis_client, "get", MagicMock(return_value=None))
vector._client.indices.exists.return_value = False
vector.create_collection([[0.1, 0.2]], [{}])
vector._client.indices.create.assert_called_once()
kwargs = vector._client.indices.create.call_args.kwargs
mappings = kwargs["mappings"]
assert mappings["properties"][huawei_module.Field.VECTOR]["dimension"] == 2
assert kwargs["settings"] == {"index.vector": True}
huawei_module.redis_client.set.assert_called_once()
def test_huawei_factory_branches(huawei_module, monkeypatch):
factory = huawei_module.HuaweiCloudVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(huawei_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(huawei_module.dify_config, "HUAWEI_CLOUD_HOSTS", "http://huawei-es:9200")
monkeypatch.setattr(huawei_module.dify_config, "HUAWEI_CLOUD_USER", "user")
monkeypatch.setattr(huawei_module.dify_config, "HUAWEI_CLOUD_PASSWORD", "pass")
with patch.object(huawei_module, "HuaweiCloudVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["index_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["index_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -1,412 +0,0 @@
import importlib
import sys
import types
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_iris_module():
iris = types.ModuleType("iris")
def connect(**_kwargs):
conn = MagicMock()
conn.cursor.return_value = MagicMock()
return conn
iris.connect = MagicMock(side_effect=connect)
return iris
@pytest.fixture
def iris_module(monkeypatch):
monkeypatch.setitem(sys.modules, "iris", _build_fake_iris_module())
import core.rag.datasource.vdb.iris.iris_vector as module
reloaded = importlib.reload(module)
reloaded._pool_instance = None
return reloaded
def _config(module, **overrides):
values = {
"IRIS_HOST": "localhost",
"IRIS_SUPER_SERVER_PORT": 1972,
"IRIS_USER": "user",
"IRIS_PASSWORD": "pass",
"IRIS_DATABASE": "db",
"IRIS_SCHEMA": "schema",
"IRIS_CONNECTION_URL": "url",
"IRIS_MIN_CONNECTION": 1,
"IRIS_MAX_CONNECTION": 2,
"IRIS_TEXT_INDEX": True,
"IRIS_TEXT_INDEX_LANGUAGE": "en",
}
values.update(overrides)
return module.IrisVectorConfig.model_validate(values)
def test_get_iris_pool_singleton(iris_module):
iris_module._pool_instance = None
cfg = _config(iris_module)
with patch.object(iris_module, "IrisConnectionPool", return_value="pool") as pool_cls:
pool_1 = iris_module.get_iris_pool(cfg)
pool_2 = iris_module.get_iris_pool(cfg)
assert pool_1 == "pool"
assert pool_2 == "pool"
pool_cls.assert_called_once_with(cfg)
@pytest.fixture
def pool_with_min_max(iris_module):
cfg = _config(iris_module, IRIS_MIN_CONNECTION=2, IRIS_MAX_CONNECTION=3)
with patch.object(iris_module.IrisConnectionPool, "_create_connection", return_value=MagicMock()) as create_conn:
pool = iris_module.IrisConnectionPool(cfg)
yield pool, create_conn
def test_pool_initialization_respects_min_max(pool_with_min_max):
pool, create_conn = pool_with_min_max
assert len(pool._pool) == 2
assert create_conn.call_count == 2
@pytest.fixture
def pool_for_get_connection(iris_module):
cfg = _config(iris_module, IRIS_MIN_CONNECTION=2, IRIS_MAX_CONNECTION=3)
pool = iris_module.IrisConnectionPool(cfg)
return pool
def test_get_connection_returns_existing_and_increments(pool_for_get_connection):
pool = pool_for_get_connection
conn = MagicMock()
pool._pool = [conn]
pool._in_use = 0
assert pool.get_connection() is conn
assert pool._in_use == 1
def test_get_connection_creates_new_when_empty(pool_for_get_connection):
pool = pool_for_get_connection
pool._pool = []
pool._in_use = 0
pool._create_connection = MagicMock(return_value="new-conn")
assert pool.get_connection() == "new-conn"
def test_get_connection_raises_when_exhausted(pool_for_get_connection):
pool = pool_for_get_connection
pool._pool = []
pool._in_use = pool._max_size
with pytest.raises(RuntimeError, match="exhausted"):
pool.get_connection()
@pytest.fixture
def pool_for_return_connection(iris_module):
cfg = _config(iris_module)
with patch.object(iris_module.IrisConnectionPool, "_initialize_pool", return_value=None):
pool = iris_module.IrisConnectionPool(cfg)
return pool
def test_return_connection_adds_healthy(pool_for_return_connection):
pool = pool_for_return_connection
pool._in_use = 1
conn = MagicMock()
cursor = MagicMock()
conn.cursor.return_value = cursor
pool.return_connection(conn)
assert pool._pool[-1] is conn
assert pool._in_use == 0
def test_return_connection_replaces_bad(pool_for_return_connection):
pool = pool_for_return_connection
pool._in_use = 1
bad_conn = MagicMock()
bad_cursor = MagicMock()
bad_cursor.execute.side_effect = OSError("bad")
bad_conn.cursor.return_value = bad_cursor
replacement = MagicMock()
pool._create_connection = MagicMock(return_value=replacement)
pool.return_connection(bad_conn)
bad_conn.close.assert_called_once()
assert pool._pool[-1] is replacement
assert pool._in_use == 0
def test_return_connection_ignores_none(pool_for_return_connection):
pool = pool_for_return_connection
before = len(pool._pool)
pool.return_connection(None)
assert len(pool._pool) == before
@pytest.fixture
def pool_for_schema_and_close(iris_module):
cfg = _config(iris_module)
with patch.object(iris_module.IrisConnectionPool, "_initialize_pool", return_value=None):
pool = iris_module.IrisConnectionPool(cfg)
conn = MagicMock()
cursor = MagicMock()
conn.cursor.return_value = cursor
pool._pool = [conn]
return pool, conn, cursor
def test_ensure_schema_exists_cached_noop(pool_for_schema_and_close):
pool, conn, cursor = pool_for_schema_and_close
pool._schemas_initialized = {"cached_schema"}
pool.ensure_schema_exists("cached_schema")
cursor.execute.assert_not_called()
def test_ensure_schema_exists_creates_new(pool_for_schema_and_close):
pool, conn, cursor = pool_for_schema_and_close
pool._schemas_initialized = set()
cursor.fetchone.return_value = (0,)
pool.ensure_schema_exists("new_schema")
assert "new_schema" in pool._schemas_initialized
assert any("CREATE SCHEMA" in call.args[0] for call in cursor.execute.call_args_list)
conn.commit.assert_called_once()
def test_ensure_schema_exists_existing_no_commit(pool_for_schema_and_close):
pool, conn, cursor = pool_for_schema_and_close
pool._schemas_initialized = set()
cursor.fetchone.return_value = (1,)
pool.ensure_schema_exists("existing_schema")
conn.commit.assert_not_called()
def test_ensure_schema_exists_rollback_on_error(pool_for_schema_and_close):
pool, conn, cursor = pool_for_schema_and_close
pool._schemas_initialized = set()
cursor.execute.side_effect = RuntimeError("schema failure")
with pytest.raises(RuntimeError, match="schema failure"):
pool.ensure_schema_exists("broken_schema")
conn.rollback.assert_called()
def test_close_all_closes_and_resets(iris_module):
cfg = _config(iris_module)
with patch.object(iris_module.IrisConnectionPool, "_initialize_pool", return_value=None):
pool = iris_module.IrisConnectionPool(cfg)
conn = MagicMock()
conn_2 = MagicMock()
conn_2.close.side_effect = OSError("close fail")
pool._pool = [conn, conn_2]
pool._schemas_initialized = {"x"}
pool.close_all()
assert pool._pool == []
assert pool._in_use == 0
assert pool._schemas_initialized == set()
def test_iris_vector_init_get_cursor_and_create(iris_module):
pool = MagicMock()
pool.get_connection.return_value = MagicMock()
with patch.object(iris_module, "get_iris_pool", return_value=pool):
vector = iris_module.IrisVector("collection", _config(iris_module))
assert vector.table_name == "EMBEDDING_COLLECTION"
assert vector.schema == "schema"
assert vector.get_type() == iris_module.VectorType.IRIS
conn = MagicMock()
cursor = MagicMock()
conn.cursor.return_value = cursor
vector.pool.get_connection.return_value = conn
with vector._get_cursor() as got_cursor:
assert got_cursor is cursor
conn.commit.assert_called_once()
vector.pool.return_connection.assert_called_with(conn)
conn = MagicMock()
cursor = MagicMock()
conn.cursor.return_value = cursor
vector.pool.get_connection.return_value = conn
with pytest.raises(RuntimeError, match="boom"):
with vector._get_cursor():
raise RuntimeError("boom")
conn.rollback.assert_called_once()
vector._create_collection = MagicMock()
vector.add_texts = MagicMock(return_value=["id-1"])
docs = [Document(page_content="a", metadata={"doc_id": "id-1"})]
assert vector.create(docs, [[0.1, 0.2]]) == ["id-1"]
vector._create_collection.assert_called_once_with(2)
def test_iris_vector_crud_and_vector_search(iris_module, monkeypatch):
with patch.object(iris_module, "get_iris_pool", return_value=MagicMock()):
vector = iris_module.IrisVector("collection", _config(iris_module))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
monkeypatch.setattr(iris_module.uuid, "uuid4", lambda: "generated-id")
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
SimpleNamespace(page_content="b", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "generated-id"]
assert cursor.execute.call_count == 2
cursor.fetchone.return_value = (1,)
assert vector.text_exists("id-1") is True
cursor.fetchone.return_value = None
assert vector.text_exists("id-2") is False
vector._get_cursor = MagicMock(side_effect=RuntimeError("db down"))
assert vector.text_exists("id-3") is False
vector._get_cursor = _cursor_ctx
vector.delete_by_ids([])
before = cursor.execute.call_count
vector.delete_by_ids(["id-1", "id-2"])
assert cursor.execute.call_count == before + 1
vector.delete_by_metadata_field("document_id", "doc-1")
assert "meta LIKE" in cursor.execute.call_args.args[0]
cursor.fetchall.return_value = [
("id-1", "text-1", '{"document_id":"d-1"}', 0.9),
("id-2", "text-2", '{"document_id":"d-2"}', 0.2),
("id-x",),
]
docs = vector.search_by_vector([0.1, 0.2], top_k=3, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
def test_iris_vector_full_text_search_paths(iris_module, monkeypatch):
cfg = _config(iris_module, IRIS_TEXT_INDEX=True)
with patch.object(iris_module, "get_iris_pool", return_value=MagicMock()):
vector = iris_module.IrisVector("collection", cfg)
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
cursor.execute.side_effect = None
cursor.fetchall.return_value = [
("id-1", "text-1", '{"document_id":"d-1"}', 0.7),
("id-2", "text-2", "{}", None),
]
docs = vector.search_by_full_text("query", top_k=2, document_ids_filter=["d-1"])
assert len(docs) == 2
assert docs[0].metadata["score"] == pytest.approx(0.7)
assert docs[1].metadata["score"] == pytest.approx(0.0)
cursor.reset_mock()
cursor.execute.side_effect = [RuntimeError("rank failed"), None]
cursor.fetchall.return_value = [("id-3", "text-3", "{}", 0.5)]
docs = vector.search_by_full_text("query", top_k=1)
assert len(docs) == 1
assert cursor.execute.call_count == 2
cfg_like = _config(iris_module, IRIS_TEXT_INDEX=False)
with patch.object(iris_module, "get_iris_pool", return_value=MagicMock()):
vector_like = iris_module.IrisVector("collection", cfg_like)
vector_like._get_cursor = _cursor_ctx
fake_libs = types.ModuleType("libs")
fake_helper = types.ModuleType("libs.helper")
fake_helper.escape_like_pattern = lambda value: value.replace("%", "\\%")
monkeypatch.setitem(sys.modules, "libs", fake_libs)
monkeypatch.setitem(sys.modules, "libs.helper", fake_helper)
cursor.reset_mock()
cursor.execute.side_effect = None
cursor.fetchall.return_value = []
assert vector_like.search_by_full_text("100%", top_k=1) == []
def test_iris_vector_delete_create_collection_and_factory(iris_module, monkeypatch):
with patch.object(iris_module, "get_iris_pool", return_value=MagicMock()):
vector = iris_module.IrisVector("collection", _config(iris_module, IRIS_TEXT_INDEX=True))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector.delete()
assert "DROP TABLE" in cursor.execute.call_args.args[0]
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(iris_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(iris_module.redis_client, "set", MagicMock())
monkeypatch.setattr(iris_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(2)
cursor.execute.assert_called_once()
cursor.reset_mock()
monkeypatch.setattr(iris_module.redis_client, "get", MagicMock(return_value=None))
vector.pool.ensure_schema_exists = MagicMock()
vector._create_collection(3)
assert cursor.execute.call_count == 3
iris_module.redis_client.set.assert_called_once()
cursor.reset_mock()
vector.config.IRIS_TEXT_INDEX = False
vector._create_collection(3)
assert cursor.execute.call_count == 2
factory = iris_module.IrisVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(iris_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(iris_module.dify_config, "IRIS_HOST", "localhost")
monkeypatch.setattr(iris_module.dify_config, "IRIS_SUPER_SERVER_PORT", 1972)
monkeypatch.setattr(iris_module.dify_config, "IRIS_USER", "user")
monkeypatch.setattr(iris_module.dify_config, "IRIS_PASSWORD", "pass")
monkeypatch.setattr(iris_module.dify_config, "IRIS_DATABASE", "db")
monkeypatch.setattr(iris_module.dify_config, "IRIS_SCHEMA", "schema")
monkeypatch.setattr(iris_module.dify_config, "IRIS_CONNECTION_URL", "url")
monkeypatch.setattr(iris_module.dify_config, "IRIS_MIN_CONNECTION", 1)
monkeypatch.setattr(iris_module.dify_config, "IRIS_MAX_CONNECTION", 2)
monkeypatch.setattr(iris_module.dify_config, "IRIS_TEXT_INDEX", True)
monkeypatch.setattr(iris_module.dify_config, "IRIS_TEXT_INDEX_LANGUAGE", "en")
with patch.object(iris_module, "IrisVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,394 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_opensearch_modules():
opensearchpy = types.ModuleType("opensearchpy")
opensearch_helpers = types.ModuleType("opensearchpy.helpers")
class BulkIndexError(Exception):
def __init__(self, errors):
super().__init__("bulk error")
self.errors = errors
class OpenSearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.indices = SimpleNamespace(
refresh=MagicMock(),
exists=MagicMock(return_value=False),
delete=MagicMock(),
create=MagicMock(),
)
self.bulk = MagicMock(return_value={"errors": False, "items": []})
self.search = MagicMock(return_value={"hits": {"hits": []}})
self.delete_by_query = MagicMock()
self.get = MagicMock(return_value={"_id": "id"})
self.exists = MagicMock(return_value=True)
opensearch_helpers.BulkIndexError = BulkIndexError
opensearch_helpers.bulk = MagicMock()
opensearchpy.OpenSearch = OpenSearch
opensearchpy.helpers = opensearch_helpers
return {
"opensearchpy": opensearchpy,
"opensearchpy.helpers": opensearch_helpers,
}
@pytest.fixture
def lindorm_module(monkeypatch):
for name, module in _build_fake_opensearch_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.lindorm.lindorm_vector as module
return importlib.reload(module)
def _config(module):
return module.LindormVectorStoreConfig(
hosts="http://localhost:9200",
username="user",
password="pass",
using_ugc=False,
request_timeout=3.0,
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("hosts", None, "config URL is required"),
("username", None, "config USERNAME is required"),
("password", None, "config PASSWORD is required"),
],
)
def test_lindorm_config_validation(lindorm_module, field, value, message):
values = _config(lindorm_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
lindorm_module.LindormVectorStoreConfig.model_validate(values)
def test_to_opensearch_params_and_init(lindorm_module):
cfg = _config(lindorm_module)
params = cfg.to_opensearch_params()
assert params["hosts"] == "http://localhost:9200"
assert params["http_auth"] == ("user", "pass")
vector = lindorm_module.LindormVectorStore("Collection", cfg, using_ugc=False)
assert vector._collection_name == "collection"
assert vector.get_type() == lindorm_module.VectorType.LINDORM
with pytest.raises(ValueError, match="routing_value"):
lindorm_module.LindormVectorStore("c", cfg, using_ugc=True)
vector_ugc = lindorm_module.LindormVectorStore("c", cfg, using_ugc=True, routing_value="ROUTE")
assert vector_ugc._routing == "route"
def test_create_refresh_and_add_texts_success(lindorm_module, monkeypatch):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="a", metadata={"doc_id": "id-1"})]
vector.create(docs, [[0.1]])
vector.create_collection.assert_called_once_with([[0.1]], [{"doc_id": "id-1"}])
vector.add_texts.assert_called_once_with(docs, [[0.1]])
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
monkeypatch.setattr(lindorm_module.time, "sleep", MagicMock())
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
Document(page_content="c", metadata={"doc_id": "id-3"}),
]
embeddings = [[0.1], [0.2], [0.3]]
vector.add_texts(docs, embeddings, batch_size=2, timeout=9)
assert vector._client.bulk.call_count == 2
actions = vector._client.bulk.call_args_list[0].args[0]
assert actions[0]["index"]["routing"] == "route"
assert actions[1][lindorm_module.ROUTING_FIELD] == "route"
vector.refresh()
vector._client.indices.refresh.assert_called_once_with(index="collection")
def test_add_texts_error_paths(lindorm_module):
vector = lindorm_module.LindormVectorStore("collection", _config(lindorm_module), using_ugc=False)
vector._client.bulk.return_value = {"errors": True, "items": [{"index": {"error": "boom"}}]}
docs = [Document(page_content="a", metadata={"doc_id": "id-1"})]
with pytest.raises(Exception, match="RetryError"):
vector.add_texts(docs, [[0.1]], batch_size=1)
vector._client.bulk.side_effect = RuntimeError("bulk failed")
with pytest.raises(Exception, match="RetryError"):
vector.add_texts(docs, [[0.1]], batch_size=1)
def test_metadata_lookup_and_delete_by_metadata(lindorm_module):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
vector._client.search.return_value = {"hits": {"hits": [{"_id": "id-1"}, {"_id": "id-2"}]}}
ids = vector.get_ids_by_metadata_field("document_id", "doc-1")
assert ids == ["id-1", "id-2"]
query = vector._client.search.call_args.kwargs["body"]
must_conditions = query["query"]["bool"]["must"]
assert any("routing_field.keyword" in cond.get("term", {}) for cond in must_conditions)
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_ids.assert_called_once_with(["id-1", "id-2"])
vector._client.search.return_value = {"hits": {"hits": []}}
vector.delete_by_ids.reset_mock()
vector.delete_by_metadata_field("document_id", "doc-2")
vector.delete_by_ids.assert_not_called()
def test_delete_by_ids_paths(lindorm_module):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
vector.delete_by_ids([])
vector._client.indices.exists.assert_not_called()
vector._client.indices.exists.return_value = False
vector.delete_by_ids(["id-1"])
vector._client.indices.exists.return_value = True
vector._client.exists.side_effect = [True, False]
lindorm_module.helpers.bulk.reset_mock()
vector.delete_by_ids(["id-1", "id-2"])
lindorm_module.helpers.bulk.assert_called_once()
actions = lindorm_module.helpers.bulk.call_args.args[1]
assert len(actions) == 1
assert actions[0]["routing"] == "route"
lindorm_module.helpers.bulk.reset_mock()
lindorm_module.helpers.bulk.side_effect = lindorm_module.BulkIndexError(
errors=[
{"delete": {"status": 404, "_id": "id-404"}},
{"delete": {"status": 500, "_id": "id-500"}},
]
)
vector._client.exists.side_effect = [True]
vector.delete_by_ids(["id-1"])
def test_delete_and_text_exists(lindorm_module):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
vector.delete()
vector._client.delete_by_query.assert_called_once()
vector._client.indices.refresh.assert_called_once_with(index="collection")
vector = lindorm_module.LindormVectorStore("collection", _config(lindorm_module), using_ugc=False)
vector._client.indices.exists.return_value = True
vector.delete()
vector._client.indices.delete.assert_called_once_with(index="collection", params={"timeout": 60})
vector._client.indices.delete.reset_mock()
vector._client.indices.exists.return_value = False
vector.delete()
vector._client.indices.delete.assert_not_called()
assert vector.text_exists("id-1") is True
vector._client.get.side_effect = RuntimeError("missing")
assert vector.text_exists("id-1") is False
def test_search_by_vector_validation_and_success(lindorm_module):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
with pytest.raises(ValueError, match="should be a list"):
vector.search_by_vector("bad")
with pytest.raises(ValueError, match="should be floats"):
vector.search_by_vector([0.1, "bad"])
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_score": 0.9,
"_source": {
lindorm_module.Field.CONTENT_KEY: "doc-a",
lindorm_module.Field.VECTOR: [0.1],
lindorm_module.Field.METADATA_KEY: {"doc_id": "1", "document_id": "d-1"},
},
},
{
"_score": 0.2,
"_source": {
lindorm_module.Field.CONTENT_KEY: "doc-b",
lindorm_module.Field.VECTOR: [0.2],
lindorm_module.Field.METADATA_KEY: {"doc_id": "2", "document_id": "d-2"},
},
},
]
}
}
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
call_kwargs = vector._client.search.call_args.kwargs
query = call_kwargs["body"]
assert "ext" in query
assert query["query"]["knn"][lindorm_module.Field.VECTOR]["filter"]["bool"]["must"]
assert call_kwargs["params"]["routing"] == "route"
vector._client.search.side_effect = RuntimeError("search failed")
with pytest.raises(RuntimeError, match="search failed"):
vector.search_by_vector([0.1])
def test_search_by_full_text_success_and_error(lindorm_module):
vector = lindorm_module.LindormVectorStore(
"collection", _config(lindorm_module), using_ugc=True, routing_value="route"
)
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_source": {
lindorm_module.Field.CONTENT_KEY: "doc-a",
lindorm_module.Field.VECTOR: [0.1],
lindorm_module.Field.METADATA_KEY: {"doc_id": "1"},
}
}
]
}
}
docs = vector.search_by_full_text("hello", top_k=2, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].page_content == "doc-a"
query = vector._client.search.call_args.kwargs["body"]
assert query["query"]["bool"]["filter"]
vector._client.search.side_effect = RuntimeError("full text failed")
with pytest.raises(RuntimeError, match="full text failed"):
vector.search_by_full_text("hello")
def test_create_collection_paths(lindorm_module, monkeypatch):
vector = lindorm_module.LindormVectorStore("collection", _config(lindorm_module), using_ugc=False)
with pytest.raises(ValueError, match="cannot be empty"):
vector.create_collection([])
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(lindorm_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(lindorm_module.redis_client, "set", MagicMock())
monkeypatch.setattr(lindorm_module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection([[0.1, 0.2]])
vector._client.indices.create.assert_not_called()
monkeypatch.setattr(lindorm_module.redis_client, "get", MagicMock(return_value=None))
vector._client.indices.exists.return_value = False
vector.create_collection([[0.1, 0.2]], index_params={"index_type": "ivf", "space_type": "cosine"})
vector._client.indices.create.assert_called_once()
body = vector._client.indices.create.call_args.kwargs["body"]
assert body["mappings"]["properties"][lindorm_module.Field.VECTOR]["method"]["name"] == "ivf"
assert body["mappings"]["properties"][lindorm_module.Field.VECTOR]["method"]["space_type"] == "cosine"
vector._client.indices.create.reset_mock()
vector._client.indices.exists.return_value = True
vector.create_collection([[0.1, 0.2]])
vector._client.indices.create.assert_not_called()
def test_lindorm_factory_branches(lindorm_module, monkeypatch):
factory = lindorm_module.LindormVectorStoreFactory()
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_URL", "http://localhost:9200")
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_USERNAME", "user")
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_PASSWORD", "pass")
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_QUERY_TIMEOUT", 3.0)
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_INDEX_TYPE", "hnsw")
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_DISTANCE_TYPE", "l2")
monkeypatch.setattr(lindorm_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
dataset = SimpleNamespace(id="dataset-1", index_struct=None, index_struct_dict={})
embeddings = SimpleNamespace(embed_query=lambda _q: [0.1, 0.2, 0.3])
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_USING_UGC", None)
with pytest.raises(ValueError, match="LINDORM_USING_UGC is not set"):
factory.init_vector(dataset, attributes=[], embeddings=embeddings)
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_USING_UGC", False)
dataset_existing_plain = SimpleNamespace(
id="dataset-1",
index_struct="{}",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING"}, "using_ugc": False},
)
with patch.object(lindorm_module, "LindormVectorStore", return_value="vector") as store_cls:
result = factory.init_vector(dataset_existing_plain, attributes=[], embeddings=embeddings)
assert result == "vector"
assert store_cls.call_args.args[0] == "existing"
dataset_existing_ugc = SimpleNamespace(
id="dataset-1",
index_struct="{}",
index_struct_dict={
"vector_store": {"class_prefix": "ROUTING"},
"using_ugc": True,
"dimension": 1536,
"index_type": "hnsw",
"distance_type": "l2",
},
)
with patch.object(lindorm_module, "LindormVectorStore", return_value="vector") as store_cls:
factory.init_vector(dataset_existing_ugc, attributes=[], embeddings=embeddings)
assert store_cls.call_args.args[0] == "ugc_index_1536_hnsw_l2"
assert store_cls.call_args.kwargs["routing_value"] == "ROUTING"
dataset_new = SimpleNamespace(id="dataset-2", index_struct=None, index_struct_dict={})
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_USING_UGC", True)
with patch.object(lindorm_module, "LindormVectorStore", return_value="vector") as store_cls:
factory.init_vector(dataset_new, attributes=[], embeddings=embeddings)
assert store_cls.call_args.args[0] == "ugc_index_3_hnsw_l2"
assert store_cls.call_args.kwargs["routing_value"] == "auto_collection"
assert dataset_new.index_struct is not None
dataset_new_plain = SimpleNamespace(id="dataset-3", index_struct=None, index_struct_dict={})
monkeypatch.setattr(lindorm_module.dify_config, "LINDORM_USING_UGC", False)
with patch.object(lindorm_module, "LindormVectorStore", return_value="vector") as store_cls:
factory.init_vector(dataset_new_plain, attributes=[], embeddings=embeddings)
assert store_cls.call_args.args[0] == "auto_collection"
assert store_cls.call_args.kwargs["routing_value"] is None
@@ -1,252 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_mo_vector_modules():
mo_vector = types.ModuleType("mo_vector")
mo_vector.__path__ = []
mo_vector_client = types.ModuleType("mo_vector.client")
class MoVectorClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.create_full_text_index = MagicMock()
self.insert = MagicMock()
self.get = MagicMock(return_value=[])
self.delete = MagicMock()
self.query_by_metadata = MagicMock(return_value=[])
self.query = MagicMock(return_value=[])
self.full_text_query = MagicMock(return_value=[])
mo_vector_client.MoVectorClient = MoVectorClient
mo_vector.client = mo_vector_client
return {"mo_vector": mo_vector, "mo_vector.client": mo_vector_client}
@pytest.fixture
def matrixone_module(monkeypatch):
for name, module in _build_fake_mo_vector_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.matrixone.matrixone_vector as module
return importlib.reload(module)
def _valid_config(module):
return module.MatrixoneConfig(
host="localhost",
port=6001,
user="dump",
password="111",
database="dify",
metric="l2",
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config host is required"),
("port", 0, "config port is required"),
("user", "", "config user is required"),
("password", "", "config password is required"),
("database", "", "config database is required"),
],
)
def test_matrixone_config_validation(matrixone_module, field, value, message):
values = _valid_config(matrixone_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
matrixone_module.MatrixoneConfig.model_validate(values)
def test_get_client_creates_full_text_index_when_cache_misses(matrixone_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(matrixone_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(matrixone_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(matrixone_module.redis_client, "set", MagicMock())
vector = matrixone_module.MatrixoneVector("Collection_1", _valid_config(matrixone_module))
client = vector._get_client(dimension=3, create_table=True)
assert client.kwargs["table_name"] == "collection_1"
client.create_full_text_index.assert_called_once()
matrixone_module.redis_client.set.assert_called_once()
def test_get_client_skips_index_creation_when_cache_hits(matrixone_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(matrixone_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(matrixone_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(matrixone_module.redis_client, "set", MagicMock())
vector = matrixone_module.MatrixoneVector("Collection_1", _valid_config(matrixone_module))
client = vector._get_client(dimension=3, create_table=True)
client.create_full_text_index.assert_not_called()
matrixone_module.redis_client.set.assert_not_called()
def test_ensure_client_initializes_client_for_decorated_methods(matrixone_module):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
vector.client = None
fake_client = MagicMock()
fake_client.get.return_value = [{"id": "seg-1"}]
vector._get_client = MagicMock(return_value=fake_client)
exists = vector.text_exists("seg-1")
assert exists is True
vector._get_client.assert_called_once_with(None, False)
def test_search_by_full_text_parses_metadata_and_applies_threshold(matrixone_module):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
vector.client = MagicMock()
vector.client.full_text_query.return_value = [
SimpleNamespace(document="doc-a", metadata='{"doc_id":"1"}', distance=0.1),
SimpleNamespace(document="doc-b", metadata={"doc_id": "2"}, distance=0.7),
]
docs = vector.search_by_full_text("query", top_k=2, score_threshold=0.5, document_ids_filter=["doc-1"])
assert len(docs) == 1
assert docs[0].page_content == "doc-a"
assert docs[0].metadata["doc_id"] == "1"
assert docs[0].metadata["score"] == pytest.approx(0.9)
assert vector.client.full_text_query.call_args.kwargs["filter"] == {"document_id": {"$in": ["doc-1"]}}
def test_get_type_and_create_delegate_to_add_texts(matrixone_module):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
fake_client = MagicMock()
vector._get_client = MagicMock(return_value=fake_client)
vector.add_texts = MagicMock(return_value=["seg-1"])
docs = [Document(page_content="hello", metadata={"doc_id": "seg-1"})]
result = vector.create(docs, [[0.1, 0.2]])
assert vector.get_type() == "matrixone"
assert result == ["seg-1"]
vector._get_client.assert_called_once_with(2, True)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_get_client_handles_full_text_index_creation_error(matrixone_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(matrixone_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(matrixone_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(matrixone_module.redis_client, "set", MagicMock())
failing_client = MagicMock()
failing_client.create_full_text_index.side_effect = RuntimeError("boom")
monkeypatch.setattr(matrixone_module, "MoVectorClient", MagicMock(return_value=failing_client))
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
client = vector._get_client(dimension=3, create_table=True)
assert client is failing_client
matrixone_module.redis_client.set.assert_not_called()
def test_add_texts_generates_ids_and_inserts(matrixone_module, monkeypatch):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
vector.client = MagicMock()
monkeypatch.setattr(matrixone_module.uuid, "uuid4", lambda: "generated-uuid")
docs = [
Document(page_content="a", metadata={"doc_id": "doc-a", "document_id": "d-1"}),
Document(page_content="b", metadata={"document_id": "d-2"}),
SimpleNamespace(page_content="c", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]])
# For current prod code, only docs with metadata get ids, so only two ids
assert ids == ["doc-a", "generated-uuid"]
vector.client.insert.assert_called_once()
insert_kwargs = vector.client.insert.call_args.kwargs
# All lists passed to insert should be the same length
texts = insert_kwargs["texts"]
embeddings = insert_kwargs["embeddings"]
metadatas = insert_kwargs["metadatas"]
ids_insert = insert_kwargs["ids"]
assert len(texts) == len(embeddings) == len(metadatas) == len(docs)
# ids may be shorter than docs for current prod code, but should match number of docs with metadata
assert ids_insert == ["doc-a", "generated-uuid"]
def test_delete_and_metadata_methods(matrixone_module):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
vector.client = MagicMock()
vector.client.query_by_metadata.return_value = [SimpleNamespace(id="seg-1"), SimpleNamespace(id="seg-2")]
vector.delete_by_ids([])
vector.client.delete.assert_not_called()
vector.delete_by_ids(["seg-1"])
vector.delete_by_metadata_field("document_id", "doc-1")
ids = vector.get_ids_by_metadata_field("document_id", "doc-1")
vector.delete()
assert ids == ["seg-1", "seg-2"]
assert vector.client.delete.call_count == 3
def test_search_by_vector_builds_documents(matrixone_module):
vector = matrixone_module.MatrixoneVector("collection_1", _valid_config(matrixone_module))
vector.client = MagicMock()
vector.client.query.return_value = [
SimpleNamespace(document="doc-a", metadata={"doc_id": "1"}),
SimpleNamespace(document="doc-b", metadata={"doc_id": "2"}),
]
docs = vector.search_by_vector([0.1, 0.2], top_k=2, document_ids_filter=["d-1"])
assert len(docs) == 2
assert docs[0].page_content == "doc-a"
assert docs[1].metadata["doc_id"] == "2"
assert vector.client.query.call_args.kwargs["filter"] == {"document_id": {"$in": ["d-1"]}}
def test_matrixone_factory_uses_existing_or_generated_collection(matrixone_module, monkeypatch):
factory = matrixone_module.MatrixoneVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(matrixone_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_HOST", "127.0.0.1")
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_PORT", 6001)
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_USER", "dump")
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_PASSWORD", "111")
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_DATABASE", "dify")
monkeypatch.setattr(matrixone_module.dify_config, "MATRIXONE_METRIC", "l2")
with patch.object(matrixone_module, "MatrixoneVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,414 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_pymilvus_modules():
pymilvus = types.ModuleType("pymilvus")
pymilvus.__path__ = []
pymilvus_milvus_client = types.ModuleType("pymilvus.milvus_client")
pymilvus_orm = types.ModuleType("pymilvus.orm")
pymilvus_orm.__path__ = []
pymilvus_orm_types = types.ModuleType("pymilvus.orm.types")
class MilvusError(Exception):
pass
class MilvusClient:
def __init__(self, **kwargs):
self.init_kwargs = kwargs
self.has_collection = MagicMock(return_value=False)
self.describe_collection = MagicMock(
return_value={"fields": [{"name": "id"}, {"name": "content"}, {"name": "metadata"}]}
)
self.get_server_version = MagicMock(return_value="2.5.0")
self.insert = MagicMock(return_value=[1])
self.query = MagicMock(return_value=[])
self.delete = MagicMock()
self.drop_collection = MagicMock()
self.search = MagicMock(return_value=[[]])
self.create_collection = MagicMock()
class IndexParams:
def __init__(self):
self.indexes = []
def add_index(self, **kwargs):
self.indexes.append(kwargs)
class DataType:
JSON = "JSON"
VARCHAR = "VARCHAR"
INT64 = "INT64"
SPARSE_FLOAT_VECTOR = "SPARSE_FLOAT_VECTOR"
FLOAT_VECTOR = "FLOAT_VECTOR"
class FieldSchema:
def __init__(self, name, dtype, **kwargs):
self.name = name
self.dtype = dtype
self.kwargs = kwargs
class CollectionSchema:
def __init__(self, fields):
self.fields = fields
self.functions = []
def add_function(self, func):
self.functions.append(func)
class FunctionType:
BM25 = "BM25"
class Function:
def __init__(self, **kwargs):
self.kwargs = kwargs
def infer_dtype_bydata(_value):
return DataType.FLOAT_VECTOR
pymilvus.MilvusException = MilvusError
pymilvus.MilvusClient = MilvusClient
pymilvus.IndexParams = IndexParams
pymilvus.CollectionSchema = CollectionSchema
pymilvus.DataType = DataType
pymilvus.FieldSchema = FieldSchema
pymilvus.Function = Function
pymilvus.FunctionType = FunctionType
pymilvus_milvus_client.IndexParams = IndexParams
pymilvus_orm.types = pymilvus_orm_types
pymilvus_orm_types.infer_dtype_bydata = infer_dtype_bydata
# Attach submodules for dotted imports
pymilvus.milvus_client = pymilvus_milvus_client
pymilvus.orm = pymilvus_orm
return {
"pymilvus": pymilvus,
"pymilvus.milvus_client": pymilvus_milvus_client,
"pymilvus.orm": pymilvus_orm,
"pymilvus.orm.types": pymilvus_orm_types,
}
@pytest.fixture
def milvus_module(monkeypatch):
for name, module in _build_fake_pymilvus_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.milvus.milvus_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"uri": "http://localhost:19530",
"user": "root",
"password": "Milvus",
"database": "default",
"enable_hybrid_search": False,
"analyzer_params": None,
}
values.update(overrides)
return module.MilvusConfig.model_validate(values)
def test_config_validation_and_defaults(milvus_module):
valid_config = {"uri": "http://localhost:19530", "user": "root", "password": "Milvus"}
for key in valid_config:
config = valid_config.copy()
del config[key]
with pytest.raises(ValidationError) as e:
milvus_module.MilvusConfig.model_validate(config)
assert e.value.errors()[0]["msg"] == f"Value error, config MILVUS_{key.upper()} is required"
config = milvus_module.MilvusConfig.model_validate(valid_config)
assert config.database == "default"
token_config = milvus_module.MilvusConfig.model_validate(
{"uri": "http://localhost:19530", "token": "token-value", "database": "db-1"}
)
assert token_config.token == "token-value"
def test_config_to_milvus_params(milvus_module):
config = _config(milvus_module, analyzer_params='{"tokenizer":"standard"}')
params = config.to_milvus_params()
assert params["uri"] == "http://localhost:19530"
assert params["db_name"] == "default"
assert params["analyzer_params"] == '{"tokenizer":"standard"}'
def test_init_client_supports_token_and_user_password(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
token_client = vector._init_client(
milvus_module.MilvusConfig.model_validate({"uri": "http://localhost:19530", "token": "abc", "database": "db"})
)
assert token_client.init_kwargs == {"uri": "http://localhost:19530", "token": "abc", "db_name": "db"}
user_client = vector._init_client(_config(milvus_module))
assert user_client.init_kwargs["uri"] == "http://localhost:19530"
assert user_client.init_kwargs["user"] == "root"
assert user_client.init_kwargs["password"] == "Milvus"
def test_init_loads_fields_when_collection_exists(milvus_module):
client = milvus_module.MilvusClient(uri="http://localhost:19530")
client.has_collection.return_value = True
client.describe_collection.return_value = {
"fields": [{"name": "id"}, {"name": "content"}, {"name": "metadata"}, {"name": "sparse_vector"}]
}
with patch.object(milvus_module.MilvusVector, "_init_client", return_value=client):
with patch.object(milvus_module.MilvusVector, "_check_hybrid_search_support", return_value=False):
vector = milvus_module.MilvusVector("collection_1", _config(milvus_module))
assert "id" not in vector._fields
assert "content" in vector._fields
def test_load_collection_fields_from_argument_and_remote(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._client = MagicMock()
vector._collection_name = "collection_1"
vector._client.describe_collection.return_value = {"fields": [{"name": "id"}, {"name": "content"}]}
vector._load_collection_fields(["id", "metadata"])
assert vector._fields == ["metadata"]
vector._load_collection_fields()
assert vector._fields == ["content"]
def test_check_hybrid_search_support_branches(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._client = MagicMock()
vector._client_config = SimpleNamespace(enable_hybrid_search=False)
assert vector._check_hybrid_search_support() is False
vector._client_config = SimpleNamespace(enable_hybrid_search=True)
vector._client.get_server_version.return_value = "Zilliz Cloud 2.4"
assert vector._check_hybrid_search_support() is True
vector._client.get_server_version.return_value = "2.5.1"
assert vector._check_hybrid_search_support() is True
vector._client.get_server_version.return_value = "2.4.9"
assert vector._check_hybrid_search_support() is False
vector._client.get_server_version.side_effect = RuntimeError("boom")
assert vector._check_hybrid_search_support() is False
def test_get_type_and_create_delegate(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [SimpleNamespace(page_content="hello", metadata=None)]
vector.create(docs, [[0.1, 0.2]])
assert vector.get_type() == "milvus"
vector.create_collection.assert_called_once()
create_args = vector.create_collection.call_args.args
assert create_args[0] == [[0.1, 0.2]]
assert create_args[1] == [{}]
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_batches_and_raises_milvus_exception(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.insert.side_effect = [["id-1"], ["id-2"]]
docs = [Document(page_content=f"text-{i}", metadata={"doc_id": f"d-{i}"}) for i in range(1001)]
embeddings = [[0.1, 0.2] for _ in range(1001)]
ids = vector.add_texts(docs, embeddings)
assert ids == ["id-1", "id-2"]
assert vector._client.insert.call_count == 2
vector._client.insert.side_effect = milvus_module.MilvusException("insert failed")
with pytest.raises(milvus_module.MilvusException):
vector.add_texts([Document(page_content="x", metadata={})], [[0.1]])
def test_get_ids_and_delete_methods(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.query.return_value = [{"id": 1}, {"id": 2}]
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == [1, 2]
vector._client.query.return_value = []
assert vector.get_ids_by_metadata_field("document_id", "doc-1") is None
vector._client.has_collection.return_value = True
vector.get_ids_by_metadata_field = MagicMock(return_value=[101, 102])
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.delete.assert_called_with(collection_name="collection_1", pks=[101, 102])
vector._client.delete.reset_mock()
vector._client.query.return_value = [{"id": 11}, {"id": 12}]
vector.delete_by_ids(["doc-a", "doc-b"])
vector._client.delete.assert_called_with(collection_name="collection_1", pks=[11, 12])
vector._client.has_collection.return_value = True
vector.delete()
vector._client.drop_collection.assert_called_once_with("collection_1", None)
def test_text_exists_and_field_exists(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._fields = ["content", "metadata"]
vector._client = MagicMock()
vector._client.has_collection.return_value = False
assert vector.text_exists("doc-1") is False
vector._client.has_collection.return_value = True
vector._client.query.return_value = [{"id": 1}]
assert vector.text_exists("doc-1") is True
vector._client.query.return_value = []
assert vector.text_exists("doc-1") is False
assert vector.field_exists("content") is True
assert vector.field_exists("unknown") is False
def test_process_search_results_and_search_methods(milvus_module):
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._fields = ["content", "metadata", "sparse_vector"]
processed = vector._process_search_results(
[
[
{"entity": {"content": "doc-1", "metadata": {"doc_id": "1"}}, "distance": 0.9},
{"entity": {"content": "doc-2", "metadata": {"doc_id": "2"}}, "distance": 0.2},
]
],
[milvus_module.Field.CONTENT_KEY, milvus_module.Field.METADATA_KEY],
score_threshold=0.5,
)
assert len(processed) == 1
assert processed[0].metadata["score"] == 0.9
vector._client.search.return_value = [[{"entity": {"content": "doc"}, "distance": 0.8}]]
vector._process_search_results = MagicMock(return_value=["doc"])
docs = vector.search_by_vector([0.1, 0.2], top_k=3, document_ids_filter=["a", "b"], score_threshold=0.1)
assert docs == ["doc"]
assert vector._client.search.call_args.kwargs["filter"] == 'metadata["document_id"] in ["a", "b"]'
vector._hybrid_search_enabled = False
assert vector.search_by_full_text("query") == []
vector._hybrid_search_enabled = True
vector._fields = []
assert vector.search_by_full_text("query") == []
vector._fields = [milvus_module.Field.SPARSE_VECTOR]
vector._process_search_results = MagicMock(return_value=["full-text-doc"])
full_text_docs = vector.search_by_full_text("query", top_k=2, document_ids_filter=["d-1"], score_threshold=0.2)
assert full_text_docs == ["full-text-doc"]
assert "document_id" in vector._client.search.call_args.kwargs["filter"]
def test_create_collection_cache_and_existing_collection(milvus_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(milvus_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(milvus_module.redis_client, "set", MagicMock())
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._consistency_level = "Session"
vector._client_config = _config(milvus_module)
vector._hybrid_search_enabled = False
vector._client = MagicMock()
monkeypatch.setattr(milvus_module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection([[0.1, 0.2]], metadatas=[{"doc_id": "1"}], index_params={"index_type": "HNSW"})
vector._client.create_collection.assert_not_called()
monkeypatch.setattr(milvus_module.redis_client, "get", MagicMock(return_value=None))
vector._client.has_collection.return_value = True
vector.create_collection([[0.1, 0.2]], metadatas=[{"doc_id": "1"}], index_params={"index_type": "HNSW"})
milvus_module.redis_client.set.assert_called()
def test_create_collection_builds_schema_and_indexes(milvus_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(milvus_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(milvus_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(milvus_module.redis_client, "set", MagicMock())
vector = milvus_module.MilvusVector.__new__(milvus_module.MilvusVector)
vector._collection_name = "collection_1"
vector._consistency_level = "Session"
vector._client = MagicMock()
vector._client.has_collection.return_value = False
vector._load_collection_fields = MagicMock()
vector._client_config = _config(milvus_module, analyzer_params='{"tokenizer":"standard"}')
vector._hybrid_search_enabled = True
vector.create_collection(
embeddings=[[0.1, 0.2]],
metadatas=[{"doc_id": "1"}],
index_params={"metric_type": "IP", "index_type": "HNSW", "params": {"M": 8}},
)
call_kwargs = vector._client.create_collection.call_args.kwargs
schema = call_kwargs["schema"]
index_params_obj = call_kwargs["index_params"]
field_names = [f.name for f in schema.fields]
assert milvus_module.Field.SPARSE_VECTOR in field_names
assert len(schema.functions) == 1
assert len(index_params_obj.indexes) == 2
assert call_kwargs["consistency_level"] == "Session"
def test_factory_initializes_milvus_vector(milvus_module, monkeypatch):
factory = milvus_module.MilvusVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(milvus_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_URI", "http://localhost:19530")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_TOKEN", "")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_USER", "root")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_PASSWORD", "Milvus")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_DATABASE", "default")
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_ENABLE_HYBRID_SEARCH", True)
monkeypatch.setattr(milvus_module.dify_config, "MILVUS_ANALYZER_PARAMS", '{"tokenizer":"standard"}')
with patch.object(milvus_module, "MilvusVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,230 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_clickhouse_connect_module():
clickhouse_connect = types.ModuleType("clickhouse_connect")
class QueryResult:
def __init__(self, rows=None, named_rows=None):
self.row_count = len(rows or [])
self.result_rows = rows or []
self._named_rows = named_rows or []
def named_results(self):
return self._named_rows
class Client:
def __init__(self):
self.command = MagicMock()
self.query = MagicMock(return_value=QueryResult())
client = Client()
def get_client(**_kwargs):
return client
clickhouse_connect.get_client = get_client
clickhouse_connect.QueryResult = QueryResult
clickhouse_connect._fake_client = client
return clickhouse_connect
@pytest.fixture
def myscale_module(monkeypatch):
fake_module = _build_fake_clickhouse_connect_module()
monkeypatch.setitem(sys.modules, "clickhouse_connect", fake_module)
import core.rag.datasource.vdb.myscale.myscale_vector as module
return importlib.reload(module)
def _config(module):
return module.MyScaleConfig(
host="localhost",
port=8123,
user="default",
password="",
database="dify",
fts_params="",
)
def test_escape_str_replaces_backslash_and_quote(myscale_module):
escaped = myscale_module.MyScaleVector.escape_str(r"text\with'special")
assert escaped == "text with special"
def test_search_raises_for_invalid_top_k(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector._search("distance(vector, [0.1, 0.2])", myscale_module.SortOrder.ASC, top_k=0)
def test_search_builds_where_clause_for_cosine_threshold(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._client.query.return_value = myscale_module.get_client().query.return_value.__class__(
named_rows=[{"text": "doc-1", "vector": [0.1, 0.2], "metadata": {"doc_id": "seg-1"}}]
)
docs = vector._search("distance(vector, [0.1, 0.2])", myscale_module.SortOrder.ASC, top_k=1, score_threshold=0.2)
assert len(docs) == 1
sql = vector._client.query.call_args.args[0]
assert "WHERE dist < 0.8" in sql
def test_delete_by_ids_short_circuits_on_empty_list(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._client.command.reset_mock()
vector.delete_by_ids([])
vector._client.command.assert_not_called()
def test_factory_initializes_lower_case_collection_name(myscale_module, monkeypatch):
factory = myscale_module.MyScaleVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(myscale_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_HOST", "localhost")
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_PORT", 8123)
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_USER", "default")
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_PASSWORD", "")
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_DATABASE", "dify")
monkeypatch.setattr(myscale_module.dify_config, "MYSCALE_FTS_PARAMS", "")
with patch.object(myscale_module, "MyScaleVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
def test_init_and_get_type_set_expected_defaults(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
assert vector.get_type() == "myscale"
assert vector._vec_order == myscale_module.SortOrder.ASC
vector._client.command.assert_called_with("SET allow_experimental_object_type=1")
def test_create_calls_create_collection_and_add_texts(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._create_collection = MagicMock()
vector.add_texts = MagicMock(return_value=["seg-1"])
docs = [Document(page_content="hello", metadata={"doc_id": "seg-1"})]
result = vector.create(docs, [[0.1, 0.2]])
assert result == ["seg-1"]
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once()
def test_create_collection_builds_expected_sql(myscale_module):
config = myscale_module.MyScaleConfig(
host="localhost",
port=8123,
user="default",
password="",
database="dify",
fts_params="tokenizer=unicode",
)
vector = myscale_module.MyScaleVector("collection_1", config)
vector._client.command.reset_mock()
vector._create_collection(3)
assert vector._client.command.call_count == 2
sql = vector._client.command.call_args_list[1].args[0]
assert "CREATE TABLE IF NOT EXISTS dify.collection_1" in sql
assert "CONSTRAINT cons_vec_len CHECK length(vector) = 3" in sql
assert "INDEX text_idx text TYPE fts('tokenizer=unicode')" in sql
def test_add_texts_inserts_rows_and_returns_ids(myscale_module, monkeypatch):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
monkeypatch.setattr(myscale_module.uuid, "uuid4", lambda: "generated-uuid")
docs = [
Document(page_content=r"te'xt\1", metadata={"doc_id": "doc-a", "document_id": "d-1"}),
Document(page_content="text-2", metadata={"document_id": "d-2"}),
SimpleNamespace(page_content="text-3", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]])
assert ids == ["doc-a", "generated-uuid"]
sql = vector._client.command.call_args.args[0]
assert "INSERT INTO dify.collection_1" in sql
assert "te xt 1" in sql
def test_text_exists_and_metadata_operations(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._client.query.return_value = SimpleNamespace(row_count=1, result_rows=[("id-1",), ("id-2",)])
assert vector.text_exists("id-1") is True
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"]
vector.delete_by_ids(["id-1", "id-2"])
vector.delete_by_metadata_field("document_id", "doc-1")
assert vector._client.command.call_count >= 2
def test_search_delegation_methods(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._search = MagicMock(return_value=["result"])
result_vector = vector.search_by_vector([0.1, 0.2], top_k=2)
result_text = vector.search_by_full_text("hello", top_k=2)
assert result_vector == ["result"]
assert result_text == ["result"]
assert vector._search.call_count == 2
def test_search_with_document_filter_and_exception(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._client.query.return_value = SimpleNamespace(
named_results=lambda: [{"text": "doc", "vector": [0.1], "metadata": {"doc_id": "1"}}]
)
docs = vector._search(
"distance(vector, [0.1])",
myscale_module.SortOrder.ASC,
top_k=2,
document_ids_filter=["doc-1", "doc-2"],
)
assert len(docs) == 1
sql = vector._client.query.call_args.args[0]
assert "metadata['document_id'] in ('doc-1', 'doc-2')" in sql
vector._client.query.side_effect = RuntimeError("boom")
assert vector._search("distance(vector, [0.1])", myscale_module.SortOrder.ASC, top_k=1) == []
def test_delete_drops_table(myscale_module):
vector = myscale_module.MyScaleVector("collection_1", _config(myscale_module))
vector._client.command.reset_mock()
vector.delete()
vector._client.command.assert_called_once_with("DROP TABLE IF EXISTS dify.collection_1")
@@ -1,553 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from sqlalchemy.exc import SQLAlchemyError
from core.rag.models.document import Document
def _build_fake_pyobvector_module():
pyobvector = types.ModuleType("pyobvector")
class VECTOR:
def __init__(self, dim):
self.dim = dim
def l2_distance(*_args, **_kwargs):
return "l2"
def cosine_distance(*_args, **_kwargs):
return "cosine"
def inner_product(*_args, **_kwargs):
return "inner_product"
class ObVecClient:
def __init__(self, **_kwargs):
self.metadata_obj = SimpleNamespace(tables={})
self.engine = MagicMock()
self.check_table_exists = MagicMock(return_value=False)
self.perform_raw_text_sql = MagicMock()
self.prepare_index_params = MagicMock()
self.create_table_with_index_params = MagicMock()
self.refresh_metadata = MagicMock()
self.insert = MagicMock()
self.refresh_index = MagicMock()
self.get = MagicMock()
self.delete = MagicMock()
self.set_ob_hnsw_ef_search = MagicMock()
self.ann_search = MagicMock(return_value=[])
self.drop_table_if_exist = MagicMock()
pyobvector.VECTOR = VECTOR
pyobvector.ObVecClient = ObVecClient
pyobvector.l2_distance = l2_distance
pyobvector.cosine_distance = cosine_distance
pyobvector.inner_product = inner_product
return pyobvector
@pytest.fixture
def oceanbase_module(monkeypatch):
monkeypatch.setitem(sys.modules, "pyobvector", _build_fake_pyobvector_module())
import core.rag.datasource.vdb.oceanbase.oceanbase_vector as module
return importlib.reload(module)
def _config(module):
return module.OceanBaseVectorConfig(
host="127.0.0.1",
port=2881,
user="root",
password="secret",
database="test",
enable_hybrid_search=True,
batch_size=10,
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config OCEANBASE_VECTOR_HOST is required"),
("port", 0, "config OCEANBASE_VECTOR_PORT is required"),
("user", "", "config OCEANBASE_VECTOR_USER is required"),
("database", "", "config OCEANBASE_VECTOR_DATABASE is required"),
],
)
def test_oceanbase_config_validation(oceanbase_module, field, value, message):
values = _config(oceanbase_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
oceanbase_module.OceanBaseVectorConfig.model_validate(values)
def test_init_rejects_invalid_collection_name(oceanbase_module):
with pytest.raises(ValueError, match="Invalid collection name"):
oceanbase_module.OceanBaseVector("invalid-name", _config(oceanbase_module))
def test_distance_to_score_for_supported_metrics(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._config = SimpleNamespace(metric_type="l2")
assert vector._distance_to_score(3.0) == pytest.approx(0.25)
vector._config = SimpleNamespace(metric_type="cosine")
assert vector._distance_to_score(0.2) == pytest.approx(0.8)
vector._config = SimpleNamespace(metric_type="inner_product")
assert vector._distance_to_score(-0.2) == pytest.approx(0.2)
def test_get_distance_func_raises_for_unknown_metric(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._config = SimpleNamespace(metric_type="manhattan")
with pytest.raises(ValueError, match="Unsupported metric_type"):
vector._get_distance_func()
def test_process_search_results_handles_json_and_score_threshold(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
rows = [
("doc-1", '{"doc_id":"1"}', 0.9),
("doc-2", "not-json", 0.8),
("doc-3", {"doc_id": "3"}, 0.3),
]
docs = vector._process_search_results(rows, score_threshold=0.5, score_key="rank")
assert len(docs) == 2
assert docs[0].metadata["doc_id"] == "1"
assert docs[0].metadata["rank"] == 0.9
assert docs[1].metadata["rank"] == 0.8
def test_search_by_vector_validates_document_id_format(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._hnsw_ef_search = -1
vector._config = SimpleNamespace(metric_type="cosine")
vector._client = MagicMock()
with pytest.raises(ValueError, match="Invalid document ID format"):
vector.search_by_vector([0.1, 0.2], document_ids_filter=["bad id"])
def test_search_by_full_text_returns_empty_when_disabled(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._hybrid_search_enabled = False
vector._collection_name = "collection_1"
assert vector.search_by_full_text("query") == []
def test_check_hybrid_search_support_uses_version_comment(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._config = SimpleNamespace(enable_hybrid_search=True)
vector._client = MagicMock()
cursor = MagicMock()
cursor.fetchone.return_value = ("OceanBase_CE 4.3.5.1 (rxxxxxxxxx) (Built Mar 18 2025)",)
vector._client.perform_raw_text_sql.return_value = cursor
assert vector._check_hybrid_search_support() is True
cursor.fetchone.return_value = ("OceanBase_CE 4.3.4.0 (rxxxxxxxxx) (Built Mar 18 2025)",)
assert vector._check_hybrid_search_support() is False
def test_init_get_type_and_field_loading(oceanbase_module):
config = _config(oceanbase_module)
config.enable_hybrid_search = False
table = SimpleNamespace(columns=[SimpleNamespace(name="id"), SimpleNamespace(name="text")])
fake_client = oceanbase_module.ObVecClient()
fake_client.check_table_exists.return_value = True
fake_client.metadata_obj.tables = {"collection_1": table}
with patch.object(oceanbase_module, "ObVecClient", return_value=fake_client):
vector = oceanbase_module.OceanBaseVector("collection_1", config)
assert vector.get_type() == "oceanbase"
assert vector.field_exists("text") is True
def test_load_collection_fields_handles_missing_table_and_exception(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._fields = []
vector._client = MagicMock()
vector._client.metadata_obj.tables = {}
vector._load_collection_fields()
assert vector._fields == []
vector._client.metadata_obj.tables = {"collection_1": MagicMock(columns=MagicMock(side_effect=RuntimeError("x")))}
vector._load_collection_fields()
assert vector._fields == []
def test_create_delegates_to_collection_and_insert(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="text", metadata={"doc_id": "1"})]
vector.create(docs, [[0.1, 0.2]])
assert vector._vec_dim == 2
vector._create_collection.assert_called_once()
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_create_collection_cache_and_existing_table_short_circuits(oceanbase_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(oceanbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(oceanbase_module.redis_client, "set", MagicMock())
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._vec_dim = 2
vector._hybrid_search_enabled = False
vector._config = SimpleNamespace(metric_type="cosine", hnsw_m=16, hnsw_ef_construction=64)
vector._client = MagicMock()
vector.delete = MagicMock()
vector._load_collection_fields = MagicMock()
monkeypatch.setattr(oceanbase_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection()
vector._client.check_table_exists.assert_not_called()
monkeypatch.setattr(oceanbase_module.redis_client, "get", MagicMock(return_value=None))
vector._client.check_table_exists.return_value = True
vector._create_collection()
vector.delete.assert_not_called()
def test_create_collection_happy_path_with_hybrid_and_index(oceanbase_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(oceanbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(oceanbase_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(oceanbase_module.redis_client, "set", MagicMock())
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_FULLTEXT_PARSER", "ik")
monkeypatch.setattr(oceanbase_module, "Column", lambda *args, **kwargs: SimpleNamespace(args=args, kwargs=kwargs))
monkeypatch.setattr(oceanbase_module, "VECTOR", lambda dim: SimpleNamespace(dim=dim))
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._vec_dim = 3
vector._hybrid_search_enabled = True
vector._config = SimpleNamespace(metric_type="cosine", hnsw_m=16, hnsw_ef_construction=64)
vector._client = MagicMock()
vector._client.check_table_exists.return_value = False
vector._client.perform_raw_text_sql.side_effect = [
[[None, None, None, None, None, None, "30"]],
None,
None,
]
index_params = MagicMock()
vector._client.prepare_index_params.return_value = index_params
vector.delete = MagicMock()
vector._load_collection_fields = MagicMock()
vector._create_collection()
vector.delete.assert_called_once()
vector._client.create_table_with_index_params.assert_called_once()
index_params.add_index.assert_called_once()
vector._client.refresh_metadata.assert_called_once_with(["collection_1"])
oceanbase_module.redis_client.set.assert_called_once()
def test_create_collection_error_paths(oceanbase_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(oceanbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(oceanbase_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(oceanbase_module, "Column", lambda *args, **kwargs: SimpleNamespace(args=args, kwargs=kwargs))
monkeypatch.setattr(oceanbase_module, "VECTOR", lambda dim: SimpleNamespace(dim=dim))
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._vec_dim = 2
vector._hybrid_search_enabled = True
vector._config = SimpleNamespace(metric_type="cosine", hnsw_m=16, hnsw_ef_construction=64)
vector._client = MagicMock()
vector._client.check_table_exists.return_value = False
vector._client.prepare_index_params.return_value = MagicMock()
vector.delete = MagicMock()
vector._load_collection_fields = MagicMock()
vector._client.perform_raw_text_sql.return_value = []
with pytest.raises(ValueError, match="ob_vector_memory_limit_percentage not found"):
vector._create_collection()
vector._client.perform_raw_text_sql.side_effect = [
[[None, None, None, None, None, None, "0"]],
RuntimeError("no privilege"),
]
with pytest.raises(Exception, match="Failed to set ob_vector_memory_limit_percentage"):
vector._create_collection()
vector._client.perform_raw_text_sql.side_effect = [[[None, None, None, None, None, None, "30"]]]
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_FULLTEXT_PARSER", "not-valid")
with pytest.raises(ValueError, match="Invalid OceanBase full-text parser"):
vector._create_collection()
def test_create_collection_fulltext_and_metadata_index_exceptions(oceanbase_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(oceanbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(oceanbase_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(oceanbase_module.redis_client, "set", MagicMock())
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_FULLTEXT_PARSER", "ik")
monkeypatch.setattr(oceanbase_module, "Column", lambda *args, **kwargs: SimpleNamespace(args=args, kwargs=kwargs))
monkeypatch.setattr(oceanbase_module, "VECTOR", lambda dim: SimpleNamespace(dim=dim))
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._vec_dim = 2
vector._hybrid_search_enabled = True
vector._config = SimpleNamespace(metric_type="cosine", hnsw_m=16, hnsw_ef_construction=64)
vector._client = MagicMock()
vector._client.check_table_exists.return_value = False
vector._client.prepare_index_params.return_value = MagicMock()
vector.delete = MagicMock()
vector._load_collection_fields = MagicMock()
vector._client.perform_raw_text_sql.side_effect = [
[[None, None, None, None, None, None, "30"]],
RuntimeError("fulltext failed"),
]
with pytest.raises(Exception, match="Failed to add fulltext index"):
vector._create_collection()
vector._hybrid_search_enabled = False
vector._client.perform_raw_text_sql.side_effect = [
[[None, None, None, None, None, None, "30"]],
SQLAlchemyError("metadata index failed"),
]
vector._create_collection()
vector._client.refresh_metadata.assert_called_once_with(["collection_1"])
def test_check_hybrid_search_support_false_and_exception(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._config = SimpleNamespace(enable_hybrid_search=False)
vector._client = MagicMock()
assert vector._check_hybrid_search_support() is False
vector._config = SimpleNamespace(enable_hybrid_search=True)
vector._client.perform_raw_text_sql.side_effect = RuntimeError("boom")
assert vector._check_hybrid_search_support() is False
def test_add_texts_batches_refresh_and_exceptions(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._config = SimpleNamespace(batch_size=2, hnsw_refresh_threshold=2)
vector._client = MagicMock()
vector._get_uuids = MagicMock(return_value=["id-1", "id-2", "id-3"])
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
Document(page_content="c", metadata={"doc_id": "id-3"}),
]
vector.add_texts(docs, [[0.1], [0.2], [0.3]])
assert vector._client.insert.call_count == 2
vector._client.refresh_index.assert_called_once()
vector._client.insert.reset_mock()
vector._client.refresh_index.reset_mock()
vector._client.insert.side_effect = RuntimeError("insert failed")
with pytest.raises(Exception, match="Failed to insert batch"):
vector.add_texts([docs[0]], [[0.1]])
vector._client.insert.side_effect = None
vector._client.insert.return_value = None
vector._client.refresh_index.side_effect = SQLAlchemyError("refresh failed")
vector._config = SimpleNamespace(batch_size=10, hnsw_refresh_threshold=1)
vector._get_uuids.return_value = ["id-1"]
vector.add_texts([docs[0]], [[0.1]])
def test_text_exists_and_delete_by_ids(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector._client.get.return_value = SimpleNamespace(rowcount=1)
assert vector.text_exists("id-1") is True
vector._client.get.side_effect = RuntimeError("boom")
with pytest.raises(Exception, match="Failed to check text existence"):
vector.text_exists("id-1")
vector.delete_by_ids([])
vector._client.delete.assert_not_called()
vector._client.delete.side_effect = None
vector.delete_by_ids(["id-1"])
vector._client.delete.assert_called_once()
vector._client.delete.side_effect = RuntimeError("boom")
with pytest.raises(Exception, match="Failed to delete documents"):
vector.delete_by_ids(["id-1"])
def test_get_ids_and_delete_by_metadata_field(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
execute_result = [("id-1",), ("id-2",)]
conn = MagicMock()
conn.__enter__.return_value = conn
conn.__exit__.return_value = None
conn.execute.return_value = execute_result
vector._client.engine.connect.return_value = conn
ids = vector.get_ids_by_metadata_field("document_id", "doc-1")
assert ids == ["id-1", "id-2"]
with pytest.raises(Exception, match="Failed to query documents by metadata field"):
vector.get_ids_by_metadata_field("bad key!", "doc-1")
vector.get_ids_by_metadata_field = MagicMock(return_value=["id-1"])
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_ids.assert_called_once_with(["id-1"])
vector.get_ids_by_metadata_field = MagicMock(return_value=[])
vector.delete_by_ids.reset_mock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_ids.assert_not_called()
def test_search_by_full_text_paths(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._hybrid_search_enabled = True
vector.field_exists = MagicMock(return_value=False)
assert vector.search_by_full_text("query") == []
vector.field_exists.return_value = True
vector._client = MagicMock()
conn = MagicMock()
tx = MagicMock()
tx.__enter__.return_value = tx
tx.__exit__.return_value = None
conn.begin.return_value = tx
conn.__enter__.return_value = conn
conn.__exit__.return_value = None
conn.execute.return_value.fetchall.return_value = [("text-1", '{"doc_id":"1"}', 0.9)]
vector._client.engine.connect.return_value = conn
docs = vector.search_by_full_text("query", top_k=2, document_ids_filter=["d-1"], score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == 0.9
with pytest.raises(Exception, match="Full-text search failed"):
vector.search_by_full_text("query", top_k=0)
def test_search_by_vector_paths(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._hnsw_ef_search = -1
vector._config = SimpleNamespace(metric_type="cosine")
vector._client = MagicMock()
vector._client.ann_search.return_value = [("doc-1", '{"doc_id":"1"}', 0.2)]
vector._process_search_results = MagicMock(return_value=["doc"])
docs = vector.search_by_vector(
[0.1, 0.2],
ef_search=10,
top_k=3,
score_threshold=0.1,
document_ids_filter=["good_id"],
)
assert docs == ["doc"]
vector._client.set_ob_hnsw_ef_search.assert_called_once_with(10)
with pytest.raises(ValueError, match="Invalid score_threshold parameter"):
vector.search_by_vector([0.1], score_threshold="x")
vector._client.ann_search.side_effect = RuntimeError("boom")
with pytest.raises(Exception, match="Vector search failed"):
vector.search_by_vector([0.1], score_threshold=0.1)
def test_get_distance_func_and_distance_to_score_errors(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._config = SimpleNamespace(metric_type="cosine")
assert vector._get_distance_func() is oceanbase_module.cosine_distance
vector._config = SimpleNamespace(metric_type="unknown")
with pytest.raises(ValueError, match="Unsupported metric_type"):
vector._distance_to_score(0.1)
def test_delete_success_and_exception(oceanbase_module):
vector = oceanbase_module.OceanBaseVector.__new__(oceanbase_module.OceanBaseVector)
vector._collection_name = "collection_1"
vector._client = MagicMock()
vector.delete()
vector._client.drop_table_if_exist.assert_called_once_with("collection_1")
vector._client.drop_table_if_exist.side_effect = RuntimeError("boom")
with pytest.raises(Exception, match="Failed to delete collection"):
vector.delete()
def test_oceanbase_factory_uses_existing_or_generated_collection(oceanbase_module, monkeypatch):
factory = oceanbase_module.OceanBaseVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(oceanbase_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_HOST", "127.0.0.1")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_PORT", 2881)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_USER", "root")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_PASSWORD", "password")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_DATABASE", "test")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_ENABLE_HYBRID_SEARCH", True)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_BATCH_SIZE", 10)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_METRIC_TYPE", "cosine")
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_HNSW_M", 16)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_HNSW_EF_CONSTRUCTION", 64)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_HNSW_EF_SEARCH", -1)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_POOL_SIZE", 5)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_VECTOR_MAX_OVERFLOW", 10)
monkeypatch.setattr(oceanbase_module.dify_config, "OCEANBASE_HNSW_REFRESH_THRESHOLD", 1000)
with patch.object(oceanbase_module, "OceanBaseVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].args[0] == "existing_collection"
assert vector_cls.call_args_list[1].args[0] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -1,400 +0,0 @@
import importlib
import sys
import types
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_psycopg2_modules():
psycopg2 = types.ModuleType("psycopg2")
psycopg2.__path__ = []
psycopg2_extras = types.ModuleType("psycopg2.extras")
psycopg2_pool = types.ModuleType("psycopg2.pool")
class SimpleConnectionPool:
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
self.getconn = MagicMock()
self.putconn = MagicMock()
psycopg2_pool.SimpleConnectionPool = SimpleConnectionPool
psycopg2_extras.execute_values = MagicMock()
psycopg2.pool = psycopg2_pool
psycopg2.extras = psycopg2_extras
return {
"psycopg2": psycopg2,
"psycopg2.pool": psycopg2_pool,
"psycopg2.extras": psycopg2_extras,
}
@pytest.fixture
def opengauss_module(monkeypatch):
for name, module in _build_fake_psycopg2_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.opengauss.opengauss as module
return importlib.reload(module)
def _config(module, *, enable_pq=False):
return module.OpenGaussConfig(
host="localhost",
port=6600,
user="postgres",
password="password",
database="dify",
min_connection=1,
max_connection=5,
enable_pq=enable_pq,
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config OPENGAUSS_HOST is required"),
("port", 0, "config OPENGAUSS_PORT is required"),
("user", "", "config OPENGAUSS_USER is required"),
("password", "", "config OPENGAUSS_PASSWORD is required"),
("database", "", "config OPENGAUSS_DATABASE is required"),
("min_connection", 0, "config OPENGAUSS_MIN_CONNECTION is required"),
("max_connection", 0, "config OPENGAUSS_MAX_CONNECTION is required"),
],
)
def test_opengauss_config_validation(opengauss_module, field, value, message):
values = _config(opengauss_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
opengauss_module.OpenGaussConfig.model_validate(values)
def test_opengauss_config_validation_rejects_min_greater_than_max(opengauss_module):
values = _config(opengauss_module).model_dump()
values["min_connection"] = 6
values["max_connection"] = 5
with pytest.raises(ValidationError, match="OPENGAUSS_MIN_CONNECTION should less than OPENGAUSS_MAX_CONNECTION"):
opengauss_module.OpenGaussConfig.model_validate(values)
def test_init_sets_table_name_and_vector_type(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module))
assert vector.table_name == "embedding_collection_1"
assert vector.get_type() == "opengauss"
assert vector.pool is pool
def test_create_index_with_pq_executes_pq_sql(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opengauss_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(opengauss_module.redis_client, "set", MagicMock())
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module, enable_pq=True))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector._create_index(1536)
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("enable_pq=on" in sql for sql in executed_sql)
assert any("SET hnsw_earlystop_threshold = 320" in sql for sql in executed_sql)
opengauss_module.redis_client.set.assert_called_once()
def test_create_index_skips_index_sql_for_large_dimension(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opengauss_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(opengauss_module.redis_client, "set", MagicMock())
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module, enable_pq=False))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector._create_index(3072)
cursor.execute.assert_not_called()
opengauss_module.redis_client.set.assert_called_once()
def test_search_by_vector_validates_top_k(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_vector([0.1, 0.2], top_k=0)
def test_delete_by_ids_short_circuits_with_empty_input(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module))
vector._get_cursor = MagicMock()
vector.delete_by_ids([])
vector._get_cursor.assert_not_called()
def test_get_cursor_closes_commits_and_returns_connection(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
pool = MagicMock()
conn = MagicMock()
cur = MagicMock()
pool.getconn.return_value = conn
conn.cursor.return_value = cur
vector.pool = pool
with vector._get_cursor() as got_cur:
assert got_cur is cur
cur.close.assert_called_once()
conn.commit.assert_called_once()
pool.putconn.assert_called_once_with(conn)
def test_create_calls_collection_insert_and_index(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
vector._create_collection = MagicMock()
vector.add_texts = MagicMock()
vector._create_index = MagicMock()
docs = [Document(page_content="text", metadata={"doc_id": "seg-1"})]
vector.create(docs, [[0.1, 0.2]])
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
vector._create_index.assert_called_once_with(2)
def test_create_index_returns_early_on_cache_hit(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opengauss_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=1))
monkeypatch.setattr(opengauss_module.redis_client, "set", MagicMock())
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module))
vector._get_cursor = MagicMock()
vector._create_index(1536)
vector._get_cursor.assert_not_called()
opengauss_module.redis_client.set.assert_not_called()
def test_create_index_without_pq_executes_standard_index_sql(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opengauss_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=None))
monkeypatch.setattr(opengauss_module.redis_client, "set", MagicMock())
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module, enable_pq=False))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector._create_index(1536)
sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("embedding_cosine_embedding_collection_1_idx" in query for query in sql)
def test_add_texts_uses_execute_values(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module))
cursor = MagicMock()
opengauss_module.psycopg2.extras.execute_values.reset_mock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
docs = [
Document(page_content="text-1", metadata={"doc_id": "seg-1", "document_id": "d-1"}),
SimpleNamespace(page_content="text-2", metadata=None),
]
monkeypatch.setattr(opengauss_module.uuid, "uuid4", lambda: "generated-uuid")
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["seg-1"]
opengauss_module.psycopg2.extras.execute_values.assert_called_once()
def test_text_exists_and_get_by_ids(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.fetchone.return_value = ("seg-1",)
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1"), ({"doc_id": "2"}, "text-2")])
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
assert vector.text_exists("seg-1") is True
docs = vector.get_by_ids(["seg-1", "seg-2"])
assert len(docs) == 2
assert docs[0].page_content == "text-1"
def test_delete_and_metadata_field_queries(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector.delete_by_ids(["seg-1", "seg-2"])
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete()
sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("DELETE FROM embedding_collection_1 WHERE id IN %s" in query for query in sql)
assert any("meta->>%s = %s" in query for query in sql)
assert any("DROP TABLE IF EXISTS embedding_collection_1" in query for query in sql)
def test_search_by_vector_and_full_text(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.__iter__.return_value = iter(
[
({"doc_id": "1"}, "text-1", 0.1),
({"doc_id": "2"}, "text-2", 0.6),
]
)
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
cursor.__iter__.return_value = iter([({"doc_id": "3"}, "full-text", 0.8)])
full_docs = vector.search_by_full_text("hello world", top_k=2)
assert len(full_docs) == 1
assert full_docs[0].page_content == "full-text"
def test_search_by_full_text_validates_top_k(opengauss_module):
vector = opengauss_module.OpenGauss.__new__(opengauss_module.OpenGauss)
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_full_text("query", top_k=0)
def test_create_collection_cache_and_create_path(opengauss_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(opengauss_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opengauss_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opengauss_module.redis_client, "set", MagicMock())
vector = opengauss_module.OpenGauss("collection_1", _config(opengauss_module))
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(1536)
cursor.execute.assert_not_called()
monkeypatch.setattr(opengauss_module.redis_client, "get", MagicMock(return_value=None))
vector._create_collection(1536)
cursor.execute.assert_called_once()
opengauss_module.redis_client.set.assert_called_once()
def test_opengauss_factory_uses_existing_or_generated_collection(opengauss_module, monkeypatch):
factory = opengauss_module.OpenGaussFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(opengauss_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_HOST", "localhost")
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_PORT", 6600)
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_USER", "postgres")
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_PASSWORD", "password")
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_DATABASE", "dify")
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_MIN_CONNECTION", 1)
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_MAX_CONNECTION", 5)
monkeypatch.setattr(opengauss_module.dify_config, "OPENGAUSS_ENABLE_PQ", False)
with patch.object(opengauss_module, "OpenGauss", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,360 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_opensearch_modules():
opensearchpy = types.ModuleType("opensearchpy")
opensearchpy_helpers = types.ModuleType("opensearchpy.helpers")
class BulkIndexError(Exception):
def __init__(self, errors):
super().__init__("bulk error")
self.errors = errors
class Urllib3AWSV4SignerAuth:
def __init__(self, credentials, region, service):
self.credentials = credentials
self.region = region
self.service = service
class Urllib3HttpConnection:
pass
class _IndicesClient:
def __init__(self):
self.exists = MagicMock(return_value=False)
self.create = MagicMock()
self.delete = MagicMock()
class OpenSearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.indices = _IndicesClient()
self.search = MagicMock(return_value={"hits": {"hits": []}})
self.get = MagicMock()
helpers = SimpleNamespace(bulk=MagicMock())
opensearchpy.OpenSearch = OpenSearch
opensearchpy.Urllib3AWSV4SignerAuth = Urllib3AWSV4SignerAuth
opensearchpy.Urllib3HttpConnection = Urllib3HttpConnection
opensearchpy.helpers = helpers
opensearchpy_helpers.BulkIndexError = BulkIndexError
return {
"opensearchpy": opensearchpy,
"opensearchpy.helpers": opensearchpy_helpers,
}
@pytest.fixture
def opensearch_module(monkeypatch):
for name, module in _build_fake_opensearch_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.opensearch.opensearch_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"host": "localhost",
"port": 9200,
"secure": True,
"verify_certs": True,
"auth_method": "basic",
"user": "admin",
"password": "secret",
}
values.update(overrides)
return module.OpenSearchConfig.model_validate(values)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config OPENSEARCH_HOST is required"),
("port", 0, "config OPENSEARCH_PORT is required"),
],
)
def test_config_validation_required_fields(opensearch_module, field, value, message):
values = _config(opensearch_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
opensearch_module.OpenSearchConfig.model_validate(values)
def test_config_validation_for_aws_auth_and_https_fields(opensearch_module):
values = {
"host": "localhost",
"port": 9200,
"secure": True,
"verify_certs": True,
"auth_method": "aws_managed_iam",
"user": "admin",
"password": "secret",
}
with pytest.raises(ValidationError, match="OPENSEARCH_AWS_REGION"):
opensearch_module.OpenSearchConfig.model_validate(values)
values = _config(opensearch_module).model_dump()
values["OPENSEARCH_SECURE"] = False
values["OPENSEARCH_VERIFY_CERTS"] = True
with pytest.raises(ValidationError, match="verify_certs=True requires secure"):
opensearch_module.OpenSearchConfig.model_validate(values)
def test_create_aws_managed_iam_auth(opensearch_module, monkeypatch):
class _Session:
def get_credentials(self):
return "creds"
boto3 = types.ModuleType("boto3")
boto3.Session = _Session
monkeypatch.setitem(sys.modules, "boto3", boto3)
config = _config(
opensearch_module,
auth_method="aws_managed_iam",
aws_region="us-east-1",
aws_service="es",
)
auth = config.create_aws_managed_iam_auth()
assert auth.credentials == "creds"
assert auth.region == "us-east-1"
assert auth.service == "es"
def test_to_opensearch_params_supports_basic_and_aws(opensearch_module):
basic_params = _config(opensearch_module).to_opensearch_params()
assert basic_params["http_auth"] == ("admin", "secret")
aws_config = _config(
opensearch_module,
auth_method="aws_managed_iam",
aws_region="us-west-2",
aws_service="es",
)
with patch.object(opensearch_module.OpenSearchConfig, "create_aws_managed_iam_auth", return_value="iam-auth"):
aws_params = aws_config.to_opensearch_params()
assert aws_params["http_auth"] == "iam-auth"
def test_init_and_create_delegate_calls(opensearch_module):
vector = opensearch_module.OpenSearchVector("Collection_1", _config(opensearch_module))
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="hello", metadata={"doc_id": "seg-1"})]
vector.create(docs, [[0.1, 0.2]])
assert vector.get_type() == "opensearch"
vector.create_collection.assert_called_once_with([[0.1, 0.2]], [{"doc_id": "seg-1"}])
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_supports_regular_and_aoss_clients(opensearch_module, monkeypatch):
vector = opensearch_module.OpenSearchVector("Collection_1", _config(opensearch_module, aws_service="es"))
docs = [
Document(page_content="a", metadata={"doc_id": "1"}),
Document(page_content="b", metadata={"doc_id": "2"}),
]
monkeypatch.setattr(opensearch_module, "uuid4", lambda: SimpleNamespace(hex="generated-id"))
opensearch_module.helpers.bulk.reset_mock()
vector.add_texts(docs, [[0.1], [0.2]])
actions = opensearch_module.helpers.bulk.call_args.kwargs["actions"]
assert len(actions) == 2
assert all("_id" in action for action in actions)
vector._client_config.aws_service = "aoss"
opensearch_module.helpers.bulk.reset_mock()
vector.add_texts(docs, [[0.3], [0.4]])
aoss_actions = opensearch_module.helpers.bulk.call_args.kwargs["actions"]
assert all("_id" not in action for action in aoss_actions)
def test_metadata_lookup_and_delete_by_metadata_field(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
vector._client.search.return_value = {"hits": {"hits": [{"_id": "id-1"}, {"_id": "id-2"}]}}
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"]
vector._client.search.return_value = {"hits": {"hits": []}}
assert vector.get_ids_by_metadata_field("document_id", "doc-1") is None
vector.get_ids_by_metadata_field = MagicMock(return_value=["id-1"])
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_ids.assert_called_once_with(["id-1"])
def test_delete_by_ids_branches_and_bulk_error_handling(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
opensearch_module.helpers.bulk.reset_mock()
vector._client.indices.exists.return_value = False
vector.delete_by_ids(["doc-1"])
opensearch_module.helpers.bulk.assert_not_called()
vector._client.indices.exists.return_value = True
vector.get_ids_by_metadata_field = MagicMock(side_effect=[["es-1"], None])
vector.delete_by_ids(["doc-1", "doc-2"])
opensearch_module.helpers.bulk.assert_called_once()
opensearch_module.helpers.bulk.reset_mock()
vector.get_ids_by_metadata_field = MagicMock(return_value=["es-404"])
opensearch_module.helpers.bulk.side_effect = opensearch_module.BulkIndexError(
[{"delete": {"status": 404, "_id": "es-404"}}]
)
vector.delete_by_ids(["doc-404"])
assert opensearch_module.helpers.bulk.call_count == 1
opensearch_module.helpers.bulk.side_effect = None
def test_delete_and_text_exists(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
vector.delete()
vector._client.indices.delete.assert_called_once_with(index="collection_1", ignore_unavailable=True)
vector._client.get.return_value = {"_id": "id-1"}
assert vector.text_exists("id-1") is True
vector._client.get.side_effect = RuntimeError("not found")
assert vector.text_exists("id-1") is False
def test_search_by_vector_validates_and_builds_documents(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
with pytest.raises(ValueError, match="query_vector should be a list"):
vector.search_by_vector("not-a-list")
with pytest.raises(ValueError, match="should be floats"):
vector.search_by_vector([0.1, 1])
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_source": {
opensearch_module.Field.CONTENT_KEY: "doc-1",
opensearch_module.Field.METADATA_KEY: None,
},
"_score": 0.9,
},
{
"_source": {
opensearch_module.Field.CONTENT_KEY: "doc-2",
opensearch_module.Field.METADATA_KEY: {"doc_id": "2"},
},
"_score": 0.1,
},
]
}
}
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].page_content == "doc-1"
assert docs[0].metadata["score"] == pytest.approx(0.9)
vector.search_by_vector([0.1, 0.2], top_k=3, document_ids_filter=["doc-a", "doc-b"])
query = vector._client.search.call_args.kwargs["body"]
assert "script_score" in query["query"]
def test_search_by_vector_reraises_client_error(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
vector._client.search.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):
vector.search_by_vector([0.1, 0.2])
def test_search_by_full_text_and_filters(opensearch_module):
vector = opensearch_module.OpenSearchVector("collection_1", _config(opensearch_module))
vector._client.search.return_value = {
"hits": {
"hits": [
{
"_source": {
opensearch_module.Field.METADATA_KEY: {"doc_id": "1"},
opensearch_module.Field.VECTOR: [0.1],
opensearch_module.Field.CONTENT_KEY: "matched text",
}
},
]
}
}
docs = vector.search_by_full_text("hello", document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].page_content == "matched text"
query = vector._client.search.call_args.kwargs["body"]
assert query["query"]["bool"]["filter"] == [{"terms": {"metadata.document_id": ["d-1"]}}]
def test_create_collection_cache_and_create_path(opensearch_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(opensearch_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(opensearch_module.redis_client, "set", MagicMock())
vector = opensearch_module.OpenSearchVector("Collection_1", _config(opensearch_module))
monkeypatch.setattr(opensearch_module.redis_client, "get", MagicMock(return_value=1))
vector._client.indices.create.reset_mock()
vector.create_collection([[0.1, 0.2]])
vector._client.indices.create.assert_not_called()
monkeypatch.setattr(opensearch_module.redis_client, "get", MagicMock(return_value=None))
vector._client.indices.exists.return_value = False
vector.create_collection([[0.1, 0.2]])
vector._client.indices.create.assert_called_once()
index_body = vector._client.indices.create.call_args.kwargs["body"]
assert index_body["mappings"]["properties"]["vector"]["dimension"] == 2
opensearch_module.redis_client.set.assert_called()
def test_opensearch_factory_initializes_expected_collection_name(opensearch_module, monkeypatch):
factory = opensearch_module.OpenSearchVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(opensearch_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_HOST", "localhost")
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_PORT", 9200)
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_SECURE", True)
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_VERIFY_CERTS", True)
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_AUTH_METHOD", "basic")
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_USER", "admin")
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_PASSWORD", "secret")
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_AWS_REGION", None)
monkeypatch.setattr(opensearch_module.dify_config, "OPENSEARCH_AWS_SERVICE", None)
with patch.object(opensearch_module, "OpenSearchVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -1,375 +0,0 @@
import array
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import numpy
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_oracle_modules():
jieba = types.ModuleType("jieba")
jieba_posseg = types.ModuleType("jieba.posseg")
jieba_posseg.cut = MagicMock(return_value=[])
jieba.posseg = jieba_posseg
oracledb = types.ModuleType("oracledb")
oracledb_connection = types.ModuleType("oracledb.connection")
class Connection:
pass
oracledb_connection.Connection = Connection
oracledb.defaults = SimpleNamespace(fetch_lobs=True)
oracledb.DB_TYPE_VECTOR = object()
oracledb.create_pool = MagicMock(return_value=MagicMock(release=MagicMock()))
oracledb.connect = MagicMock()
return {
"jieba": jieba,
"jieba.posseg": jieba_posseg,
"oracledb": oracledb,
"oracledb.connection": oracledb_connection,
}
def _connection_with_cursor(cursor):
cursor_ctx = MagicMock()
cursor_ctx.__enter__.return_value = cursor
cursor_ctx.__exit__.return_value = None
connection = MagicMock()
connection.__enter__.return_value = connection
connection.__exit__.return_value = None
connection.cursor.return_value = cursor_ctx
return connection
@pytest.fixture
def oracle_module(monkeypatch):
for name, module in _build_fake_oracle_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.oracle.oraclevector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"user": "system",
"password": "oracle",
"dsn": "oracle:1521/freepdb1",
"is_autonomous": False,
}
values.update(overrides)
return module.OracleVectorConfig.model_validate(values)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("user", "", "config ORACLE_USER is required"),
("password", "", "config ORACLE_PASSWORD is required"),
("dsn", "", "config ORACLE_DSN is required"),
],
)
def test_oracle_config_validation_required_fields(oracle_module, field, value, message):
values = _config(oracle_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
oracle_module.OracleVectorConfig.model_validate(values)
def test_oracle_config_validation_autonomous_requirements(oracle_module):
with pytest.raises(ValidationError, match="config_dir is required"):
oracle_module.OracleVectorConfig.model_validate(
{"user": "u", "password": "p", "dsn": "d", "is_autonomous": True}
)
def test_init_and_get_type(oracle_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(oracle_module.oracledb, "create_pool", MagicMock(return_value=pool))
vector = oracle_module.OracleVector("collection_1", _config(oracle_module))
assert vector.get_type() == "oracle"
assert vector.table_name == "embedding_collection_1"
assert vector.pool is pool
def test_numpy_converters_and_type_handlers(oracle_module):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
in_float64 = vector.numpy_converter_in(numpy.array([0.1], dtype=numpy.float64))
in_float32 = vector.numpy_converter_in(numpy.array([0.1], dtype=numpy.float32))
in_int8 = vector.numpy_converter_in(numpy.array([1], dtype=numpy.int8))
assert in_float64.typecode == "d"
assert in_float32.typecode == "f"
assert in_int8.typecode == "b"
cursor = MagicMock()
vector.input_type_handler(cursor, numpy.array([0.1], dtype=numpy.float32), 2)
cursor.var.assert_called_with(
oracle_module.oracledb.DB_TYPE_VECTOR,
arraysize=2,
inconverter=vector.numpy_converter_in,
)
metadata = SimpleNamespace(type_code=oracle_module.oracledb.DB_TYPE_VECTOR)
cursor.arraysize = 3
vector.output_type_handler(cursor, metadata)
cursor.var.assert_called_with(
metadata.type_code,
arraysize=3,
outconverter=vector.numpy_converter_out,
)
out_int8 = vector.numpy_converter_out(array.array("b", [1]))
assert out_int8.dtype == numpy.int8
out_float32 = vector.numpy_converter_out(array.array("f", [1.0]))
assert out_float32.dtype == numpy.float32
out_float64 = vector.numpy_converter_out(array.array("d", [1.0]))
assert out_float64.dtype == numpy.float64
def test_get_connection_supports_standard_and_autonomous_paths(oracle_module, monkeypatch):
connect = MagicMock(return_value="connection")
monkeypatch.setattr(oracle_module.oracledb, "connect", connect)
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.config = _config(oracle_module)
assert vector._get_connection() == "connection"
connect.assert_called_with(user="system", password="oracle", dsn="oracle:1521/freepdb1")
vector.config = _config(
oracle_module,
is_autonomous=True,
config_dir="/wallet",
wallet_location="/wallet",
wallet_password="pw",
)
vector._get_connection()
assert connect.call_args.kwargs["config_dir"] == "/wallet"
assert connect.call_args.kwargs["wallet_location"] == "/wallet"
def test_create_delegates_collection_and_insert(oracle_module):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector._create_collection = MagicMock()
vector.add_texts = MagicMock(return_value=["seg-1"])
docs = [Document(page_content="doc", metadata={"doc_id": "seg-1"})]
result = vector.create(docs, [[0.1, 0.2]])
assert result == ["seg-1"]
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_inserts_and_logs_on_failures(oracle_module, monkeypatch):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
vector.input_type_handler = MagicMock()
vector.output_type_handler = MagicMock()
cursor = MagicMock()
cursor.execute.side_effect = [None, RuntimeError("insert failed")]
connection = _connection_with_cursor(cursor)
vector._get_connection = MagicMock(return_value=connection)
monkeypatch.setattr(oracle_module.uuid, "uuid4", lambda: "generated-uuid")
docs = [
Document(page_content="a", metadata={"doc_id": "doc-a"}),
Document(page_content="b", metadata={"document_id": "doc-b"}),
SimpleNamespace(page_content="c", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]])
assert ids == ["doc-a", "generated-uuid"]
assert cursor.execute.call_count == 2
assert connection.commit.call_count >= 1
connection.close.assert_called()
def test_text_exists_and_get_by_ids(oracle_module):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
vector.pool = MagicMock()
cursor = MagicMock()
cursor.fetchone.return_value = ("id-1",)
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1"), ({"doc_id": "2"}, "text-2")])
vector._get_connection = MagicMock(return_value=_connection_with_cursor(cursor))
assert vector.text_exists("id-1") is True
docs = vector.get_by_ids(["id-1", "id-2"])
assert len(docs) == 2
assert docs[0].page_content == "text-1"
vector.pool.release.assert_called_once()
assert vector.get_by_ids([]) == []
def test_delete_methods(oracle_module):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
vector._get_connection = MagicMock(return_value=_connection_with_cursor(cursor))
vector.delete_by_ids([])
vector._get_connection.assert_not_called()
vector.delete_by_ids(["id-1", "id-2"])
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete()
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("DELETE FROM embedding_collection_1 WHERE id IN" in sql for sql in executed_sql)
assert any("JSON_VALUE(meta" in sql for sql in executed_sql)
assert any("DROP TABLE IF EXISTS embedding_collection_1" in sql for sql in executed_sql)
def test_search_by_vector_with_threshold_and_filter(oracle_module):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
vector.input_type_handler = MagicMock()
vector.output_type_handler = MagicMock()
cursor = MagicMock()
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "doc-1", 0.1), ({"doc_id": "2"}, "doc-2", 0.8)])
connection = _connection_with_cursor(cursor)
vector._get_connection = MagicMock(return_value=connection)
docs = vector.search_by_vector(
[0.1, 0.2],
top_k=0,
score_threshold=0.5,
document_ids_filter=["d-1", "d-2"],
)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
sql = cursor.execute.call_args.args[0]
assert "fetch first 4 rows only" in sql
assert "JSON_VALUE(meta, '$.document_id') IN (:2, :3)" in sql
def _fake_nltk_module(*, missing_data=False):
nltk = types.ModuleType("nltk")
nltk_corpus = types.ModuleType("nltk.corpus")
class _Data:
@staticmethod
def find(_path):
if missing_data:
raise LookupError("missing")
return True
nltk.data = _Data()
nltk.word_tokenize = lambda text: text.split()
nltk_corpus.stopwords = SimpleNamespace(words=lambda _lang: ["and", "the"])
return nltk, nltk_corpus
def test_search_by_full_text_chinese_and_english_paths(oracle_module, monkeypatch):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1", [0.1, 0.2])])
vector._get_connection = MagicMock(return_value=_connection_with_cursor(cursor))
monkeypatch.setattr(oracle_module.pseg, "cut", MagicMock(return_value=[("张", "nr"), ("三", "nr"), ("。", "x")]))
zh_docs = vector.search_by_full_text("张三", top_k=2)
assert len(zh_docs) == 1
zh_params = cursor.execute.call_args.args[1]
assert zh_params["kk"] == "张三"
nltk, nltk_corpus = _fake_nltk_module(missing_data=False)
monkeypatch.setitem(sys.modules, "nltk", nltk)
monkeypatch.setitem(sys.modules, "nltk.corpus", nltk_corpus)
cursor.__iter__.return_value = iter([({"doc_id": "2"}, "text-2", [0.3, 0.4])])
en_docs = vector.search_by_full_text("alice and bob", top_k=-1, document_ids_filter=["d-1"])
assert len(en_docs) == 1
en_sql = cursor.execute.call_args.args[0]
en_params = cursor.execute.call_args.args[1]
assert "fetch first 5 rows only" in en_sql
assert "doc_id_0" in en_params
def test_search_by_full_text_empty_query_and_missing_nltk(oracle_module, monkeypatch):
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector.table_name = "embedding_collection_1"
vector._get_connection = MagicMock()
empty_result = vector.search_by_full_text("")
assert empty_result[0].page_content == ""
nltk, nltk_corpus = _fake_nltk_module(missing_data=True)
monkeypatch.setitem(sys.modules, "nltk", nltk)
monkeypatch.setitem(sys.modules, "nltk.corpus", nltk_corpus)
with pytest.raises(LookupError, match="required NLTK data package"):
vector.search_by_full_text("english query")
def test_create_collection_cache_and_execute_path(oracle_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(oracle_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(oracle_module.redis_client, "set", MagicMock())
vector = oracle_module.OracleVector.__new__(oracle_module.OracleVector)
vector._collection_name = "collection_1"
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
vector._get_connection = MagicMock(return_value=_connection_with_cursor(cursor))
monkeypatch.setattr(oracle_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(2)
cursor.execute.assert_not_called()
monkeypatch.setattr(oracle_module.redis_client, "get", MagicMock(return_value=None))
vector._create_collection(2)
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("CREATE TABLE IF NOT EXISTS embedding_collection_1" in sql for sql in executed_sql)
assert any("CREATE INDEX IF NOT EXISTS idx_docs_embedding_collection_1" in sql for sql in executed_sql)
oracle_module.redis_client.set.assert_called_once()
def test_oracle_factory_init_vector_uses_existing_or_generated_collection(oracle_module, monkeypatch):
factory = oracle_module.OracleVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(oracle_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_USER", "system")
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_PASSWORD", "oracle")
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_DSN", "oracle:1521/freepdb1")
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_CONFIG_DIR", None)
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_WALLET_LOCATION", None)
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_WALLET_PASSWORD", None)
monkeypatch.setattr(oracle_module.dify_config, "ORACLE_IS_AUTONOMOUS", False)
with patch.object(oracle_module, "OracleVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,344 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from sqlalchemy.types import UserDefinedType
from core.rag.models.document import Document
def _build_fake_pgvecto_modules():
pgvecto_rs = types.ModuleType("pgvecto_rs")
pgvecto_rs_sqlalchemy = types.ModuleType("pgvecto_rs.sqlalchemy")
class VECTOR(UserDefinedType):
def __init__(self, dim):
self.dim = dim
pgvecto_rs_sqlalchemy.VECTOR = VECTOR
return {
"pgvecto_rs": pgvecto_rs,
"pgvecto_rs.sqlalchemy": pgvecto_rs_sqlalchemy,
}
class _FakeSessionContext:
def __init__(self, calls, execute_results=None):
self.calls = calls
self.execute_results = execute_results or []
self.execute = MagicMock(side_effect=self._execute_side_effect)
self.commit = MagicMock()
def _execute_side_effect(self, *args, **kwargs):
self.calls.append((args, kwargs))
if self.execute_results:
return self.execute_results.pop(0)
return MagicMock()
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return None
def _session_factory(calls, execute_results=None):
def _session(_client):
return _FakeSessionContext(calls=calls, execute_results=execute_results)
return _session
class _FakeBeginContext:
def __init__(self, session):
self._session = session
def __enter__(self):
return self._session
def __exit__(self, exc_type, exc, tb):
return None
def _sessionmaker_factory(calls, execute_results=None):
def _sessionmaker(*args, **kwargs):
session = _FakeSessionContext(calls=calls, execute_results=execute_results)
return MagicMock(begin=MagicMock(return_value=_FakeBeginContext(session)))
return _sessionmaker
def _patch_both(monkeypatch, module, calls, execute_results=None):
"""Patch both Session and sessionmaker on the module with the same call tracker."""
monkeypatch.setattr(module, "Session", _session_factory(calls, execute_results))
monkeypatch.setattr(module, "sessionmaker", _sessionmaker_factory(calls, execute_results))
@pytest.fixture
def pgvecto_module(monkeypatch):
for name, module in _build_fake_pgvecto_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.pgvecto_rs.collection as collection_module
import core.rag.datasource.vdb.pgvecto_rs.pgvecto_rs as module
return importlib.reload(module), importlib.reload(collection_module)
def _config(module, **overrides):
values = {
"host": "localhost",
"port": 5432,
"user": "postgres",
"password": "secret",
"database": "postgres",
}
values.update(overrides)
return module.PgvectoRSConfig.model_validate(values)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config PGVECTO_RS_HOST is required"),
("port", 0, "config PGVECTO_RS_PORT is required"),
("user", "", "config PGVECTO_RS_USER is required"),
("password", "", "config PGVECTO_RS_PASSWORD is required"),
("database", "", "config PGVECTO_RS_DATABASE is required"),
],
)
def test_pgvecto_config_validation(pgvecto_module, field, value, message):
module, _ = pgvecto_module
values = _config(module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
module.PgvectoRSConfig.model_validate(values)
def test_collection_base_has_expected_annotations(pgvecto_module):
_, collection_module = pgvecto_module
annotations = collection_module.CollectionORM.__annotations__
assert {"id", "text", "meta", "vector"} <= set(annotations)
def test_init_get_type_and_create_delegate(pgvecto_module, monkeypatch):
module, _ = pgvecto_module
session_calls = []
monkeypatch.setattr(module, "create_engine", MagicMock(return_value="engine"))
_patch_both(monkeypatch, module, session_calls)
vector = module.PGVectoRS("collection_1", _config(module), dim=3)
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="hello", metadata={"doc_id": "1"})]
vector.create(docs, [[0.1, 0.2]])
assert vector.get_type() == module.VectorType.PGVECTO_RS
module.create_engine.assert_called_once_with("postgresql+psycopg2://postgres:secret@localhost:5432/postgres")
assert any("CREATE EXTENSION IF NOT EXISTS vectors" in str(args[0]) for args, _ in session_calls)
vector.create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_create_collection_cache_and_sql_execution(pgvecto_module, monkeypatch):
module, _ = pgvecto_module
session_calls = []
monkeypatch.setattr(module, "create_engine", MagicMock(return_value="engine"))
_patch_both(monkeypatch, module, session_calls)
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(module.redis_client, "set", MagicMock())
vector = module.PGVectoRS("collection_1", _config(module), dim=3)
monkeypatch.setattr(module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection(3)
assert not any("CREATE TABLE IF NOT EXISTS collection_1" in str(args[0]) for args, _ in session_calls)
monkeypatch.setattr(module.redis_client, "get", MagicMock(return_value=None))
vector.create_collection(3)
assert any("CREATE TABLE IF NOT EXISTS collection_1" in str(args[0]) for args, _ in session_calls)
assert any("CREATE INDEX IF NOT EXISTS collection_1_embedding_index" in str(args[0]) for args, _ in session_calls)
module.redis_client.set.assert_called()
def test_add_texts_get_ids_and_delete_methods(pgvecto_module, monkeypatch):
module, _ = pgvecto_module
init_calls = []
runtime_calls = []
execute_results = [SimpleNamespace(fetchall=lambda: [("id-1",), ("id-2",)]), SimpleNamespace(fetchall=lambda: [])]
monkeypatch.setattr(module, "create_engine", MagicMock(return_value="engine"))
_patch_both(monkeypatch, module, init_calls)
vector = module.PGVectoRS("collection_1", _config(module), dim=3)
_patch_both(monkeypatch, module, runtime_calls, execute_results=list(execute_results))
class _InsertBuilder:
def __init__(self, table):
self.table = table
def values(self, **kwargs):
return ("insert", kwargs)
monkeypatch.setattr(module, "insert", lambda table: _InsertBuilder(table))
monkeypatch.setattr(module, "uuid4", MagicMock(side_effect=["uuid-1", "uuid-2"]))
docs = [
Document(page_content="a", metadata={"doc_id": "1"}),
Document(page_content="b", metadata={"doc_id": "2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["uuid-1", "uuid-2"]
assert any(call[0][0][0] == "insert" for call in runtime_calls if call[0])
monkeypatch.setattr(
module,
"Session",
_session_factory(runtime_calls, execute_results=[SimpleNamespace(fetchall=lambda: [("id-1",), ("id-2",)])]),
)
monkeypatch.setattr(module, "sessionmaker", _sessionmaker_factory(runtime_calls))
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-1", "id-2"]
monkeypatch.setattr(
module,
"Session",
_session_factory(runtime_calls, execute_results=[SimpleNamespace(fetchall=lambda: [])]),
)
assert vector.get_ids_by_metadata_field("document_id", "doc-1") is None
vector.get_ids_by_metadata_field = MagicMock(return_value=["id-1"])
vector.delete_by_metadata_field("document_id", "doc-1")
assert any("DELETE FROM collection_1 WHERE id = ANY(:ids)" in str(args[0]) for args, _ in runtime_calls)
runtime_calls.clear()
monkeypatch.setattr(
module,
"Session",
_session_factory(
runtime_calls,
execute_results=[
SimpleNamespace(fetchall=lambda: [("row-id-1",)]),
MagicMock(),
],
),
)
monkeypatch.setattr(module, "sessionmaker", _sessionmaker_factory(runtime_calls))
vector.delete_by_ids(["doc-1"])
assert any("meta->>'doc_id' = ANY (:doc_ids)" in str(args[0]) for args, _ in runtime_calls)
assert any("DELETE FROM collection_1 WHERE id = ANY(:ids)" in str(args[0]) for args, _ in runtime_calls)
runtime_calls.clear()
_patch_both(monkeypatch, module, runtime_calls, execute_results=[MagicMock()])
vector.delete()
assert any("DROP TABLE IF EXISTS collection_1" in str(args[0]) for args, _ in runtime_calls)
def test_text_exists_search_and_full_text(pgvecto_module, monkeypatch):
module, _ = pgvecto_module
init_calls = []
monkeypatch.setattr(module, "create_engine", MagicMock(return_value="engine"))
_patch_both(monkeypatch, module, init_calls)
vector = module.PGVectoRS("collection_1", _config(module), dim=3)
runtime_calls = []
monkeypatch.setattr(
module,
"Session",
_session_factory(
runtime_calls,
execute_results=[
SimpleNamespace(fetchall=lambda: [("id-1",)]),
SimpleNamespace(fetchall=lambda: []),
],
),
)
assert vector.text_exists("doc-1") is True
assert vector.text_exists("doc-1") is False
class _DistanceExpr:
def label(self, _name):
return self
class _VectorColumn:
def op(self, _operator, return_type=None):
def _call(_query_vector):
return _DistanceExpr()
return _call
class _MetaFilter:
def in_(self, values):
return ("in", values)
class _MetaColumn:
def __getitem__(self, _item):
return _MetaFilter()
class _Stmt:
def __init__(self):
self.where_called = False
def limit(self, _value):
return self
def order_by(self, _value):
return self
def where(self, _value):
self.where_called = True
return self
stmt = _Stmt()
monkeypatch.setattr(module, "select", lambda *_args: stmt)
vector._table = SimpleNamespace(vector=_VectorColumn(), meta=_MetaColumn())
rows = [
(SimpleNamespace(meta={"doc_id": "1"}, text="text-1"), 0.1),
(SimpleNamespace(meta={"doc_id": "2"}, text="text-2"), 0.8),
]
_patch_both(monkeypatch, module, runtime_calls, execute_results=[rows])
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
assert stmt.where_called is True
assert vector.search_by_full_text("hello") == []
def test_factory_uses_existing_or_generated_collection(pgvecto_module, monkeypatch):
module, _ = pgvecto_module
factory = module.PGVectoRSFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(module.dify_config, "PGVECTO_RS_HOST", "localhost")
monkeypatch.setattr(module.dify_config, "PGVECTO_RS_PORT", 5432)
monkeypatch.setattr(module.dify_config, "PGVECTO_RS_USER", "postgres")
monkeypatch.setattr(module.dify_config, "PGVECTO_RS_PASSWORD", "secret")
monkeypatch.setattr(module.dify_config, "PGVECTO_RS_DATABASE", "postgres")
embeddings = MagicMock()
embeddings.embed_query.return_value = [0.1, 0.2, 0.3]
with patch.object(module, "PGVectoRS", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=embeddings)
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=embeddings)
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -1,497 +0,0 @@
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import core.rag.datasource.vdb.pgvector.pgvector as pgvector_module
from core.rag.datasource.vdb.pgvector.pgvector import (
PGVector,
PGVectorConfig,
)
from core.rag.models.document import Document
class TestPGVector:
def setup_method(self, method):
self.config = PGVectorConfig(
host="localhost",
port=5432,
user="test_user",
password="test_password",
database="test_db",
min_connection=1,
max_connection=5,
pg_bigm=False,
)
self.collection_name = "test_collection"
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
def test_init(self, mock_pool_class):
"""Test PGVector initialization."""
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
pgvector = PGVector(self.collection_name, self.config)
assert pgvector._collection_name == self.collection_name
assert pgvector.table_name == f"embedding_{self.collection_name}"
assert pgvector.get_type() == "pgvector"
assert pgvector.pool is not None
assert pgvector.pg_bigm is False
assert pgvector.index_hash is not None
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
def test_init_with_pg_bigm(self, mock_pool_class):
"""Test PGVector initialization with pg_bigm enabled."""
config = PGVectorConfig(
host="localhost",
port=5432,
user="test_user",
password="test_password",
database="test_db",
min_connection=1,
max_connection=5,
pg_bigm=True,
)
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
pgvector = PGVector(self.collection_name, config)
assert pgvector.pg_bigm is True
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_basic(self, mock_redis, mock_pool_class):
"""Test basic collection creation."""
# Mock Redis operations
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.return_value = [1] # vector extension exists
pgvector = PGVector(self.collection_name, self.config)
pgvector._create_collection(1536)
# Verify SQL execution calls
assert mock_cursor.execute.called
# Check that CREATE TABLE was called with correct dimension
create_table_calls = [call for call in mock_cursor.execute.call_args_list if "CREATE TABLE" in str(call)]
assert len(create_table_calls) == 1
assert "vector(1536)" in create_table_calls[0][0][0]
# Check that CREATE INDEX was called (dimension <= 2000)
create_index_calls = [
call for call in mock_cursor.execute.call_args_list if "CREATE INDEX" in str(call) and "hnsw" in str(call)
]
assert len(create_index_calls) == 1
# Verify Redis cache was set
mock_redis.set.assert_called_once()
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_with_large_dimension(self, mock_redis, mock_pool_class):
"""Test collection creation with dimension > 2000 (no HNSW index)."""
# Mock Redis operations
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.return_value = [1] # vector extension exists
pgvector = PGVector(self.collection_name, self.config)
pgvector._create_collection(3072) # Dimension > 2000
# Check that CREATE TABLE was called
create_table_calls = [call for call in mock_cursor.execute.call_args_list if "CREATE TABLE" in str(call)]
assert len(create_table_calls) == 1
assert "vector(3072)" in create_table_calls[0][0][0]
# Check that HNSW index was NOT created (dimension > 2000)
hnsw_index_calls = [call for call in mock_cursor.execute.call_args_list if "hnsw" in str(call)]
assert len(hnsw_index_calls) == 0
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_with_pg_bigm(self, mock_redis, mock_pool_class):
"""Test collection creation with pg_bigm enabled."""
config = PGVectorConfig(
host="localhost",
port=5432,
user="test_user",
password="test_password",
database="test_db",
min_connection=1,
max_connection=5,
pg_bigm=True,
)
# Mock Redis operations
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.return_value = [1] # vector extension exists
pgvector = PGVector(self.collection_name, config)
pgvector._create_collection(1536)
# Check that pg_bigm index was created
bigm_index_calls = [call for call in mock_cursor.execute.call_args_list if "gin_bigm_ops" in str(call)]
assert len(bigm_index_calls) == 1
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_creates_vector_extension(self, mock_redis, mock_pool_class):
"""Test that vector extension is created if it doesn't exist."""
# Mock Redis operations
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
# First call: vector extension doesn't exist
mock_cursor.fetchone.return_value = None
pgvector = PGVector(self.collection_name, self.config)
pgvector._create_collection(1536)
# Check that CREATE EXTENSION was called
create_extension_calls = [
call for call in mock_cursor.execute.call_args_list if "CREATE EXTENSION vector" in str(call)
]
assert len(create_extension_calls) == 1
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_with_cache_hit(self, mock_redis, mock_pool_class):
"""Test that collection creation is skipped when cache exists."""
# Mock Redis operations - cache exists
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = 1 # Cache exists
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
pgvector = PGVector(self.collection_name, self.config)
pgvector._create_collection(1536)
# Check that no SQL was executed (early return due to cache)
assert mock_cursor.execute.call_count == 0
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
@patch("core.rag.datasource.vdb.pgvector.pgvector.redis_client")
def test_create_collection_with_redis_lock(self, mock_redis, mock_pool_class):
"""Test that Redis lock is used during collection creation."""
# Mock Redis operations
mock_lock = MagicMock()
mock_lock.__enter__ = MagicMock()
mock_lock.__exit__ = MagicMock()
mock_redis.lock.return_value = mock_lock
mock_redis.get.return_value = None
mock_redis.set.return_value = None
# Mock the connection pool
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
# Mock connection and cursor
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
mock_cursor.fetchone.return_value = [1] # vector extension exists
pgvector = PGVector(self.collection_name, self.config)
pgvector._create_collection(1536)
# Verify Redis lock was acquired with correct lock name
mock_redis.lock.assert_called_once_with("vector_indexing_test_collection_lock", timeout=20)
# Verify lock context manager was entered and exited
mock_lock.__enter__.assert_called_once()
mock_lock.__exit__.assert_called_once()
@patch("core.rag.datasource.vdb.pgvector.pgvector.psycopg2.pool.SimpleConnectionPool")
def test_get_cursor_context_manager(self, mock_pool_class):
"""Test that _get_cursor properly manages connection lifecycle."""
mock_pool = MagicMock()
mock_pool_class.return_value = mock_pool
mock_conn = MagicMock()
mock_cursor = MagicMock()
mock_pool.getconn.return_value = mock_conn
mock_conn.cursor.return_value = mock_cursor
pgvector = PGVector(self.collection_name, self.config)
with pgvector._get_cursor() as cur:
assert cur == mock_cursor
# Verify connection lifecycle methods were called
mock_pool.getconn.assert_called_once()
mock_cursor.close.assert_called_once()
mock_conn.commit.assert_called_once()
mock_pool.putconn.assert_called_once_with(mock_conn)
@pytest.mark.parametrize(
"invalid_config_override",
[
{"host": ""}, # Test empty host
{"port": 0}, # Test invalid port
{"user": ""}, # Test empty user
{"password": ""}, # Test empty password
{"database": ""}, # Test empty database
{"min_connection": 0}, # Test invalid min_connection
{"max_connection": 0}, # Test invalid max_connection
{"min_connection": 10, "max_connection": 5}, # Test min > max
],
)
def test_config_validation_parametrized(invalid_config_override):
"""Test configuration validation for various invalid inputs using parametrize."""
config = {
"host": "localhost",
"port": 5432,
"user": "test_user",
"password": "test_password",
"database": "test_db",
"min_connection": 1,
"max_connection": 5,
}
config.update(invalid_config_override)
with pytest.raises(ValueError):
PGVectorConfig(**config)
def test_create_delegates_collection_creation_and_insert():
vector = PGVector.__new__(PGVector)
vector._create_collection = MagicMock()
vector.add_texts = MagicMock(return_value=["doc-a"])
docs = [Document(page_content="hello", metadata={"doc_id": "doc-a"})]
result = vector.create(docs, [[0.1, 0.2]])
assert result == ["doc-a"]
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_add_texts_uses_execute_values_and_returns_ids(monkeypatch):
vector = PGVector.__new__(PGVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
monkeypatch.setattr(pgvector_module.uuid, "uuid4", lambda: "generated-uuid")
execute_values = MagicMock()
monkeypatch.setattr(pgvector_module.psycopg2.extras, "execute_values", execute_values)
docs = [
Document(page_content="a", metadata={"doc_id": "doc-a"}),
Document(page_content="b", metadata={"document_id": "doc-b"}),
SimpleNamespace(page_content="c", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]])
assert ids == ["doc-a", "generated-uuid"]
execute_values.assert_called_once()
def test_text_get_and_delete_methods():
vector = PGVector.__new__(PGVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.fetchone.return_value = ("id-1",)
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1"), ({"doc_id": "2"}, "text-2")])
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
assert vector.text_exists("id-1") is True
docs = vector.get_by_ids(["id-1", "id-2"])
assert len(docs) == 2
assert docs[0].page_content == "text-1"
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete()
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("meta->>%s = %s" in sql for sql in executed_sql)
assert any("DROP TABLE IF EXISTS embedding_collection_1" in sql for sql in executed_sql)
def test_delete_by_ids_handles_empty_undefined_table_and_generic_exception(monkeypatch):
vector = PGVector.__new__(PGVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
vector.delete_by_ids([])
cursor.execute.assert_not_called()
class _UndefinedTableError(Exception):
pass
monkeypatch.setattr(pgvector_module.psycopg2.errors, "UndefinedTable", _UndefinedTableError)
cursor.execute.side_effect = _UndefinedTableError("missing")
vector.delete_by_ids(["doc-1"])
cursor.execute.side_effect = RuntimeError("boom")
with pytest.raises(RuntimeError, match="boom"):
vector.delete_by_ids(["doc-1"])
def test_search_by_vector_supports_filter_and_threshold():
vector = PGVector.__new__(PGVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1", 0.1), ({"doc_id": "2"}, "text-2", 0.8)])
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_vector([0.1], top_k=0)
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
sql = cursor.execute.call_args.args[0]
assert "meta->>'document_id' in ('d-1')" in sql
def test_search_by_full_text_branches_for_bigm_and_standard():
vector = PGVector.__new__(PGVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1", 0.7)])
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_full_text("hello", top_k=0)
vector.pg_bigm = False
docs = vector.search_by_full_text("hello world", top_k=2, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.7)
standard_sql = cursor.execute.call_args.args[0]
assert "to_tsvector(text) @@ plainto_tsquery(%s)" in standard_sql
cursor.execute.reset_mock()
cursor.__iter__.return_value = iter([({"doc_id": "2"}, "text-2", 0.6)])
vector.pg_bigm = True
vector.search_by_full_text("hello world", top_k=2, document_ids_filter=["d-2"])
assert "SET pg_bigm.similarity_limit TO 0.000001" in cursor.execute.call_args_list[0].args[0]
assert "bigm_similarity" in cursor.execute.call_args_list[1].args[0]
def test_pgvector_factory_initializes_expected_collection_name(monkeypatch):
factory = pgvector_module.PGVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(pgvector_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_HOST", "localhost")
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_PORT", 5432)
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_USER", "postgres")
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_PASSWORD", "secret")
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_DATABASE", "postgres")
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_MIN_CONNECTION", 1)
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_MAX_CONNECTION", 5)
monkeypatch.setattr(pgvector_module.dify_config, "PGVECTOR_PG_BIGM", False)
with patch.object(pgvector_module, "PGVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,269 +0,0 @@
import importlib
import sys
import types
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_psycopg2_modules():
psycopg2 = types.ModuleType("psycopg2")
psycopg2.__path__ = []
psycopg2_extras = types.ModuleType("psycopg2.extras")
psycopg2_pool = types.ModuleType("psycopg2.pool")
class SimpleConnectionPool:
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
self.getconn = MagicMock()
self.putconn = MagicMock()
psycopg2_pool.SimpleConnectionPool = SimpleConnectionPool
psycopg2_extras.execute_values = MagicMock()
psycopg2.pool = psycopg2_pool
psycopg2.extras = psycopg2_extras
return {
"psycopg2": psycopg2,
"psycopg2.pool": psycopg2_pool,
"psycopg2.extras": psycopg2_extras,
}
@pytest.fixture
def vastbase_module(monkeypatch):
for name, module in _build_fake_psycopg2_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.pyvastbase.vastbase_vector as module
return importlib.reload(module)
def _config(module):
return module.VastbaseVectorConfig(
host="localhost",
port=5432,
user="dify",
password="secret",
database="dify",
min_connection=1,
max_connection=5,
)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config VASTBASE_HOST is required"),
("port", 0, "config VASTBASE_PORT is required"),
("user", "", "config VASTBASE_USER is required"),
("password", "", "config VASTBASE_PASSWORD is required"),
("database", "", "config VASTBASE_DATABASE is required"),
("min_connection", 0, "config VASTBASE_MIN_CONNECTION is required"),
("max_connection", 0, "config VASTBASE_MAX_CONNECTION is required"),
],
)
def test_vastbase_config_validation(vastbase_module, field, value, message):
values = _config(vastbase_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
vastbase_module.VastbaseVectorConfig.model_validate(values)
def test_vastbase_config_rejects_invalid_connection_window(vastbase_module):
with pytest.raises(ValidationError, match="VASTBASE_MIN_CONNECTION should less than VASTBASE_MAX_CONNECTION"):
vastbase_module.VastbaseVectorConfig.model_validate(
{
"host": "localhost",
"port": 5432,
"user": "dify",
"password": "secret",
"database": "dify",
"min_connection": 6,
"max_connection": 5,
}
)
def test_init_and_get_cursor_context_manager(vastbase_module, monkeypatch):
pool = MagicMock()
monkeypatch.setattr(vastbase_module.psycopg2.pool, "SimpleConnectionPool", MagicMock(return_value=pool))
conn = MagicMock()
cur = MagicMock()
pool.getconn.return_value = conn
conn.cursor.return_value = cur
vector = vastbase_module.VastbaseVector("collection_1", _config(vastbase_module))
assert vector.get_type() == "vastbase"
assert vector.table_name == "embedding_collection_1"
with vector._get_cursor() as got_cur:
assert got_cur is cur
cur.close.assert_called_once()
conn.commit.assert_called_once()
pool.putconn.assert_called_once_with(conn)
def test_create_and_add_texts(vastbase_module, monkeypatch):
vector = vastbase_module.VastbaseVector.__new__(vastbase_module.VastbaseVector)
vector.table_name = "embedding_collection_1"
vector._create_collection = MagicMock()
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
monkeypatch.setattr(vastbase_module.uuid, "uuid4", lambda: "generated-uuid")
docs = [
Document(page_content="a", metadata={"doc_id": "doc-a"}),
Document(page_content="b", metadata={"document_id": "doc-b"}),
SimpleNamespace(page_content="c", metadata=None),
]
ids = vector.add_texts(docs, [[0.1], [0.2], [0.3]])
assert ids == ["doc-a", "generated-uuid"]
vastbase_module.psycopg2.extras.execute_values.assert_called_once()
vector.add_texts = MagicMock(return_value=["doc-a"])
result = vector.create(docs, [[0.1], [0.2], [0.3]])
vector._create_collection.assert_called_once_with(1)
assert result == ["doc-a"]
def test_text_get_delete_and_metadata_methods(vastbase_module):
vector = vastbase_module.VastbaseVector.__new__(vastbase_module.VastbaseVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.fetchone.return_value = ("id-1",)
cursor.__iter__.return_value = iter([({"doc_id": "1"}, "text-1"), ({"doc_id": "2"}, "text-2")])
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
assert vector.text_exists("id-1") is True
docs = vector.get_by_ids(["id-1", "id-2"])
assert len(docs) == 2
assert docs[0].page_content == "text-1"
vector.delete_by_ids([])
vector.delete_by_ids(["id-1"])
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete()
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("DELETE FROM embedding_collection_1 WHERE id IN %s" in sql for sql in executed_sql)
assert any("meta->>%s = %s" in sql for sql in executed_sql)
assert any("DROP TABLE IF EXISTS embedding_collection_1" in sql for sql in executed_sql)
def test_search_by_vector_and_full_text(vastbase_module):
vector = vastbase_module.VastbaseVector.__new__(vastbase_module.VastbaseVector)
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
cursor.__iter__.return_value = iter(
[
({"doc_id": "1"}, "text-1", 0.1),
({"doc_id": "2"}, "text-2", 0.8),
]
)
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_vector([0.1, 0.2], top_k=0)
docs = vector.search_by_vector([0.1, 0.2], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
with pytest.raises(ValueError, match="top_k must be a positive integer"):
vector.search_by_full_text("hello", top_k=0)
cursor.__iter__.return_value = iter([({"doc_id": "3"}, "full-text", 0.7)])
full_docs = vector.search_by_full_text("hello world", top_k=2)
assert len(full_docs) == 1
assert full_docs[0].page_content == "full-text"
def test_create_collection_cache_and_dimension_branches(vastbase_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(vastbase_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(vastbase_module.redis_client, "set", MagicMock())
vector = vastbase_module.VastbaseVector.__new__(vastbase_module.VastbaseVector)
vector._collection_name = "collection_1"
vector.table_name = "embedding_collection_1"
cursor = MagicMock()
@contextmanager
def _cursor_ctx():
yield cursor
vector._get_cursor = _cursor_ctx
monkeypatch.setattr(vastbase_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(3)
cursor.execute.assert_not_called()
monkeypatch.setattr(vastbase_module.redis_client, "get", MagicMock(return_value=None))
vector._create_collection(17000)
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("CREATE TABLE IF NOT EXISTS embedding_collection_1" in sql for sql in executed_sql)
assert all("embedding_cosine_v1_idx" not in sql for sql in executed_sql)
cursor.execute.reset_mock()
vector._create_collection(3)
executed_sql = [call.args[0] for call in cursor.execute.call_args_list]
assert any("embedding_cosine_v1_idx" in sql for sql in executed_sql)
vastbase_module.redis_client.set.assert_called()
def test_vastbase_factory_uses_existing_or_generated_collection(vastbase_module, monkeypatch):
factory = vastbase_module.VastbaseVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(vastbase_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_HOST", "localhost")
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_PORT", 5432)
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_USER", "dify")
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_PASSWORD", "secret")
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_DATABASE", "dify")
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_MIN_CONNECTION", 1)
monkeypatch.setattr(vastbase_module.dify_config, "VASTBASE_MAX_CONNECTION", 5)
with patch.object(vastbase_module, "VastbaseVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,328 +0,0 @@
import importlib
import os
import sys
import types
from collections import UserDict
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_qdrant_modules():
qdrant_client = types.ModuleType("qdrant_client")
qdrant_http = types.ModuleType("qdrant_client.http")
qdrant_http_models = types.ModuleType("qdrant_client.http.models")
qdrant_http_exceptions = types.ModuleType("qdrant_client.http.exceptions")
qdrant_local_pkg = types.ModuleType("qdrant_client.local")
qdrant_local_mod = types.ModuleType("qdrant_client.local.qdrant_local")
class UnexpectedResponseError(Exception):
def __init__(self, status_code):
super().__init__(f"status={status_code}")
self.status_code = status_code
class FilterSelector:
def __init__(self, filter):
self.filter = filter
class HnswConfigDiff:
def __init__(self, **kwargs):
self.kwargs = kwargs
class TextIndexParams:
def __init__(self, **kwargs):
self.kwargs = kwargs
class VectorParams:
def __init__(self, **kwargs):
self.kwargs = kwargs
class PointStruct:
def __init__(self, **kwargs):
self.id = kwargs["id"]
self.vector = kwargs["vector"]
self.payload = kwargs["payload"]
class Filter:
def __init__(self, must=None):
self.must = must or []
class FieldCondition:
def __init__(self, key, match):
self.key = key
self.match = match
class MatchValue:
def __init__(self, value):
self.value = value
class MatchAny:
def __init__(self, any):
self.any = any
class MatchText:
def __init__(self, text):
self.text = text
class _Distance(UserDict):
def __getitem__(self, key):
return key
class QdrantClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.get_collections = MagicMock(return_value=SimpleNamespace(collections=[]))
self.create_collection = MagicMock()
self.create_payload_index = MagicMock()
self.upsert = MagicMock()
self.delete = MagicMock()
self.delete_collection = MagicMock()
self.retrieve = MagicMock(return_value=[])
self.search = MagicMock(return_value=[])
self.scroll = MagicMock(return_value=([], None))
class QdrantLocal(QdrantClient):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._load = MagicMock()
qdrant_client.QdrantClient = QdrantClient
qdrant_http_models.FilterSelector = FilterSelector
qdrant_http_models.HnswConfigDiff = HnswConfigDiff
qdrant_http_models.PayloadSchemaType = SimpleNamespace(KEYWORD="KEYWORD")
qdrant_http_models.TextIndexParams = TextIndexParams
qdrant_http_models.TextIndexType = SimpleNamespace(TEXT="TEXT")
qdrant_http_models.TokenizerType = SimpleNamespace(MULTILINGUAL="MULTILINGUAL")
qdrant_http_models.VectorParams = VectorParams
qdrant_http_models.Distance = _Distance()
qdrant_http_models.PointStruct = PointStruct
qdrant_http_models.Filter = Filter
qdrant_http_models.FieldCondition = FieldCondition
qdrant_http_models.MatchValue = MatchValue
qdrant_http_models.MatchAny = MatchAny
qdrant_http_models.MatchText = MatchText
qdrant_http_exceptions.UnexpectedResponse = UnexpectedResponseError
qdrant_http.models = qdrant_http_models
qdrant_local_mod.QdrantLocal = QdrantLocal
qdrant_local_pkg.qdrant_local = qdrant_local_mod
return {
"qdrant_client": qdrant_client,
"qdrant_client.http": qdrant_http,
"qdrant_client.http.models": qdrant_http_models,
"qdrant_client.http.exceptions": qdrant_http_exceptions,
"qdrant_client.local": qdrant_local_pkg,
"qdrant_client.local.qdrant_local": qdrant_local_mod,
}
@pytest.fixture
def qdrant_module(monkeypatch):
for name, module in _build_fake_qdrant_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.qdrant.qdrant_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"endpoint": "http://localhost:6333",
"api_key": "api-key",
"timeout": 20,
"root_path": "/tmp",
"grpc_port": 6334,
"prefer_grpc": False,
"replication_factor": 1,
"write_consistency_factor": 1,
}
values.update(overrides)
return module.QdrantConfig.model_validate(values)
def test_qdrant_config_to_params(qdrant_module):
url_params = _config(qdrant_module).to_qdrant_params().model_dump()
assert url_params["url"] == "http://localhost:6333"
assert url_params["verify"] is False
path_config = _config(qdrant_module, endpoint="path:storage")
assert path_config.to_qdrant_params().path == os.path.join("/tmp", "storage")
with pytest.raises(ValueError, match="Root path is not set"):
_config(qdrant_module, endpoint="path:storage", root_path=None).to_qdrant_params()
def test_init_and_basic_behaviour(qdrant_module):
vector = qdrant_module.QdrantVector("collection_1", "group-1", _config(qdrant_module))
assert vector.get_type() == qdrant_module.VectorType.QDRANT
assert vector.to_index_struct()["vector_store"]["class_prefix"] == "collection_1"
docs = [Document(page_content="a", metadata={"doc_id": "a"})]
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
vector.create(docs, [[0.1]])
vector.create_collection.assert_called_once_with("collection_1", 1)
vector.add_texts.assert_called_once()
def test_create_collection_and_add_texts(qdrant_module, monkeypatch):
vector = qdrant_module.QdrantVector("collection_1", "group-1", _config(qdrant_module))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(qdrant_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(qdrant_module.redis_client, "set", MagicMock())
monkeypatch.setattr(qdrant_module.redis_client, "get", MagicMock(return_value=1))
vector.create_collection("collection_1", 3)
vector._client.create_collection.assert_not_called()
monkeypatch.setattr(qdrant_module.redis_client, "get", MagicMock(return_value=None))
vector._client.get_collections.return_value = SimpleNamespace(collections=[])
vector.create_collection("collection_1", 3)
vector._client.create_collection.assert_called_once()
assert vector._client.create_payload_index.call_count == 4
qdrant_module.redis_client.set.assert_called_once()
# add_texts and generated batches
docs = [
Document(page_content="a", metadata={"doc_id": "id-1"}),
Document(page_content="b", metadata={"doc_id": "id-2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "id-2"]
assert vector._client.upsert.call_count == 1
payloads = qdrant_module.QdrantVector._build_payloads(
["a"], [{"doc_id": "id-1"}], "content", "metadata", "g1", "group_id"
)
assert payloads[0]["group_id"] == "g1"
with pytest.raises(ValueError, match="At least one of the texts is None"):
qdrant_module.QdrantVector._build_payloads(
[None], [{"doc_id": "id-1"}], "content", "metadata", "g1", "group_id"
)
def test_delete_and_exists_paths(qdrant_module):
vector = qdrant_module.QdrantVector("collection_1", "group-1", _config(qdrant_module))
unexpected = sys.modules["qdrant_client.http.exceptions"].UnexpectedResponse
vector._client.delete.side_effect = unexpected(404)
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.delete.side_effect = None
vector._client.delete.side_effect = unexpected(500)
with pytest.raises(unexpected):
vector.delete_by_metadata_field("document_id", "doc-1")
vector._client.delete.side_effect = None
vector._client.delete.side_effect = unexpected(404)
vector.delete()
vector._client.delete.side_effect = unexpected(500)
with pytest.raises(unexpected):
vector.delete()
vector._client.delete.side_effect = None
vector._client.delete.side_effect = unexpected(404)
vector.delete_by_ids(["doc-1"])
vector._client.delete.side_effect = unexpected(500)
with pytest.raises(unexpected):
vector.delete_by_ids(["doc-1"])
vector._client.delete.side_effect = None
vector._client.get_collections.return_value = SimpleNamespace(collections=[SimpleNamespace(name="other")])
assert vector.text_exists("id-1") is False
vector._client.get_collections.return_value = SimpleNamespace(collections=[SimpleNamespace(name="collection_1")])
vector._client.retrieve.return_value = [{"id": "id-1"}]
assert vector.text_exists("id-1") is True
def test_search_and_helper_methods(qdrant_module):
vector = qdrant_module.QdrantVector("collection_1", "group-1", _config(qdrant_module))
assert vector.search_by_vector([0.1], score_threshold=1.0) == []
vector._client.search.return_value = [
SimpleNamespace(payload=None, score=0.9, vector=[0.1]),
SimpleNamespace(payload={"metadata": {"doc_id": "1"}, "page_content": "doc-a"}, score=0.8, vector=[0.1]),
]
docs = vector.search_by_vector([0.1], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"])
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.8)
# full text search: keyword split, dedup and top_k limit
scroll_results = [
(
[
SimpleNamespace(id="p1", payload={"page_content": "doc-1", "metadata": {"doc_id": "1"}}, vector=[0.1]),
SimpleNamespace(id="p2", payload={"page_content": "doc-2", "metadata": {"doc_id": "2"}}, vector=[0.2]),
],
None,
),
(
[
SimpleNamespace(id="p2", payload={"page_content": "doc-2", "metadata": {"doc_id": "2"}}, vector=[0.2]),
],
None,
),
]
vector._client.scroll.side_effect = scroll_results
docs = vector.search_by_full_text("hello world", top_k=2, document_ids_filter=["d-1"])
assert len(docs) == 2
assert vector.search_by_full_text(" ", top_k=2) == []
local_client = qdrant_module.QdrantLocal()
vector._client = local_client
vector._reload_if_needed()
local_client._load.assert_called_once()
doc = vector._document_from_scored_point(
SimpleNamespace(payload={"page_content": "doc", "metadata": {"doc_id": "1"}}, vector=[0.1]),
"page_content",
"metadata",
)
assert doc.page_content == "doc"
def test_qdrant_factory_paths(qdrant_module, monkeypatch):
factory = qdrant_module.QdrantVectorFactory()
dataset = SimpleNamespace(
id="dataset-1",
tenant_id="tenant-1",
collection_binding_id=None,
index_struct_dict=None,
index_struct=None,
)
monkeypatch.setattr(qdrant_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(qdrant_module, "current_app", SimpleNamespace(config=SimpleNamespace(root_path="/root")))
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_URL", "http://localhost:6333")
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_API_KEY", "api-key")
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_CLIENT_TIMEOUT", 20)
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_GRPC_PORT", 6334)
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_GRPC_ENABLED", False)
monkeypatch.setattr(qdrant_module.dify_config, "QDRANT_REPLICATION_FACTOR", 1)
with patch.object(qdrant_module, "QdrantVector", return_value="vector") as vector_cls:
result = factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert result == "vector"
assert vector_cls.call_args.kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset.index_struct is not None
# collection binding lookup path
dataset.collection_binding_id = "binding-1"
dataset.index_struct_dict = {"vector_store": {"class_prefix": "existing"}}
monkeypatch.setattr(qdrant_module, "select", lambda _model: SimpleNamespace(where=lambda *_args: "stmt"))
qdrant_module.db.session.scalars = MagicMock(
return_value=SimpleNamespace(one_or_none=lambda: SimpleNamespace(collection_name="BOUND_COLLECTION"))
)
with patch.object(qdrant_module, "QdrantVector", return_value="vector") as vector_cls:
factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
assert vector_cls.call_args.kwargs["collection_name"] == "BOUND_COLLECTION"
qdrant_module.db.session.scalars = MagicMock(return_value=SimpleNamespace(one_or_none=lambda: None))
with pytest.raises(ValueError, match="Dataset Collection Bindings does not exist"):
factory.init_vector(dataset, attributes=[], embeddings=MagicMock())
@@ -1,322 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from sqlalchemy.types import UserDefinedType
from core.rag.models.document import Document
def _build_fake_relyt_modules():
pgvecto_rs = types.ModuleType("pgvecto_rs")
pgvecto_rs_sqlalchemy = types.ModuleType("pgvecto_rs.sqlalchemy")
class VECTOR(UserDefinedType):
def __init__(self, dim):
self.dim = dim
pgvecto_rs_sqlalchemy.VECTOR = VECTOR
return {
"pgvecto_rs": pgvecto_rs,
"pgvecto_rs.sqlalchemy": pgvecto_rs_sqlalchemy,
}
class _FakeSession:
def __init__(self, execute_result=None):
self.execute_result = execute_result or MagicMock(fetchall=lambda: [])
self.execute = MagicMock(return_value=self.execute_result)
self.commit = MagicMock()
def __enter__(self):
return self
def __exit__(self, exc_type, exc, tb):
return None
class _FakeBeginContext:
def __init__(self, session):
self._session = session
def __enter__(self):
return self._session
def __exit__(self, exc_type, exc, tb):
return None
def _patch_both(monkeypatch, module, session):
"""Patch both Session and sessionmaker on the module."""
monkeypatch.setattr(module, "Session", lambda _client: session)
monkeypatch.setattr(
module, "sessionmaker", lambda **kwargs: MagicMock(begin=MagicMock(return_value=_FakeBeginContext(session)))
)
@pytest.fixture
def relyt_module(monkeypatch):
for name, module in _build_fake_relyt_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.relyt.relyt_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"host": "localhost",
"port": 5432,
"user": "postgres",
"password": "secret",
"database": "relyt",
}
values.update(overrides)
return module.RelytConfig.model_validate(values)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("host", "", "config RELYT_HOST is required"),
("port", 0, "config RELYT_PORT is required"),
("user", "", "config RELYT_USER is required"),
("password", "", "config RELYT_PASSWORD is required"),
("database", "", "config RELYT_DATABASE is required"),
],
)
def test_relyt_config_validation(relyt_module, field, value, message):
values = _config(relyt_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
relyt_module.RelytConfig.model_validate(values)
def test_init_get_type_and_create_delegate(relyt_module, monkeypatch):
engine = MagicMock()
monkeypatch.setattr(relyt_module, "create_engine", MagicMock(return_value=engine))
vector = relyt_module.RelytVector("collection_1", _config(relyt_module), group_id="group-1")
vector.create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="hello", metadata={"doc_id": "seg-1"})]
vector.create(docs, [[0.1, 0.2]])
assert vector.get_type() == relyt_module.VectorType.RELYT
assert vector._url == "postgresql+psycopg2://postgres:secret@localhost:5432/relyt"
assert vector.embedding_dimension == 2
vector.create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(docs, [[0.1, 0.2]])
def test_create_collection_cache_and_sql_execution(relyt_module, monkeypatch):
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(relyt_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(relyt_module.redis_client, "set", MagicMock())
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
monkeypatch.setattr(relyt_module.redis_client, "get", MagicMock(return_value=1))
session = _FakeSession()
_patch_both(monkeypatch, relyt_module, session)
vector.create_collection(3)
session.execute.assert_not_called()
monkeypatch.setattr(relyt_module.redis_client, "get", MagicMock(return_value=None))
session = _FakeSession()
_patch_both(monkeypatch, relyt_module, session)
vector.create_collection(3)
executed_sql = [str(call.args[0]) for call in session.execute.call_args_list]
assert any("DROP TABLE IF EXISTS" in sql for sql in executed_sql)
assert any("CREATE TABLE IF NOT EXISTS" in sql for sql in executed_sql)
assert any("CREATE INDEX" in sql for sql in executed_sql)
relyt_module.redis_client.set.assert_called_once()
def test_add_texts_and_metadata_queries(relyt_module, monkeypatch):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector._group_id = "group-1"
vector.client = MagicMock()
begin_ctx = MagicMock()
begin_ctx.__enter__.return_value = None
begin_ctx.__exit__.return_value = None
conn = MagicMock()
conn.__enter__.return_value = conn
conn.__exit__.return_value = None
conn.begin.return_value = begin_ctx
vector.client.connect.return_value = conn
monkeypatch.setattr(relyt_module.uuid, "uuid1", MagicMock(side_effect=["id-1", "id-2"]))
docs = [
Document(page_content="a", metadata={"doc_id": "d-1"}),
Document(page_content="b", metadata={"doc_id": "d-2"}),
]
ids = vector.add_texts(docs, [[0.1], [0.2]])
assert ids == ["id-1", "id-2"]
assert conn.execute.call_count >= 1
first_insert_values = conn.execute.call_args.args[0].compile().params
assert "group_id" in str(first_insert_values)
session = _FakeSession(execute_result=MagicMock(fetchall=lambda: [("id-a",), ("id-b",)]))
monkeypatch.setattr(relyt_module, "Session", lambda _client: session)
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-a", "id-b"]
session = _FakeSession(execute_result=MagicMock(fetchall=lambda: []))
monkeypatch.setattr(relyt_module, "Session", lambda _client: session)
assert vector.get_ids_by_metadata_field("document_id", "doc-1") is None
# 1. delete_by_uuids: success and connect error
def test_delete_by_uuids_success_and_connect_error(relyt_module):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
with pytest.raises(ValueError, match="No ids provided"):
vector.delete_by_uuids(None)
conn = MagicMock()
conn.__enter__.return_value = conn
conn.__exit__.return_value = None
begin_ctx = MagicMock()
begin_ctx.__enter__.return_value = None
begin_ctx.__exit__.return_value = None
conn.begin.return_value = begin_ctx
vector.client.connect.return_value = conn
assert vector.delete_by_uuids(["id-1"]) is True
vector.client.connect.side_effect = RuntimeError("boom")
assert vector.delete_by_uuids(["id-1"]) is False
# 2. delete_by_metadata_field calls delete_by_uuids
def test_delete_by_metadata_field_calls_delete_by_uuids(relyt_module):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
vector.get_ids_by_metadata_field = MagicMock(return_value=["id-1"])
vector.delete_by_uuids = MagicMock(return_value=True)
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_uuids.assert_called_once_with(["id-1"])
# 3. delete_by_ids translates to uuids
def test_delete_by_ids_translates_to_uuids(relyt_module, monkeypatch):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
session = _FakeSession(execute_result=MagicMock(fetchall=lambda: [("uuid-1",), ("uuid-2",)]))
monkeypatch.setattr(relyt_module, "Session", lambda _client: session)
vector.delete_by_uuids = MagicMock(return_value=True)
vector.delete_by_ids(["doc-1", "doc-2"])
vector.delete_by_uuids.assert_called_once_with(["uuid-1", "uuid-2"])
# 4. text_exists True
def test_text_exists_true(relyt_module, monkeypatch):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
session = _FakeSession(execute_result=MagicMock(fetchall=lambda: [("id-1",)]))
monkeypatch.setattr(relyt_module, "Session", lambda _client: session)
assert vector.text_exists("doc-1") is True
# 5. text_exists False
def test_text_exists_false(relyt_module, monkeypatch):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
session = _FakeSession(execute_result=MagicMock(fetchall=lambda: []))
monkeypatch.setattr(relyt_module, "Session", lambda _client: session)
assert vector.text_exists("doc-1") is False
# 6. similarity_search_with_score_by_vector returns Documents and scores
def test_similarity_search_with_score_by_vector(relyt_module):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
result_rows = [
SimpleNamespace(document="doc-a", metadata={"doc_id": "1"}, distance=0.1),
SimpleNamespace(document="doc-b", metadata={"doc_id": "2"}, distance=0.8),
]
conn = MagicMock()
conn.__enter__.return_value = conn
conn.__exit__.return_value = None
conn.execute.return_value.fetchall.return_value = result_rows
vector.client.connect.return_value = conn
similarities = vector.similarity_search_with_score_by_vector([0.1, 0.2], k=2, filter={"document_id": ["d-1"]})
assert len(similarities) == 2
assert similarities[0][0].page_content == "doc-a"
# 7. search_by_vector filters by score and ids
def test_search_by_vector_filters_by_score_and_ids(relyt_module):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
vector.similarity_search_with_score_by_vector = MagicMock(
return_value=[
(Document(page_content="a", metadata={"doc_id": "1"}), 0.1),
(Document(page_content="b", metadata={}), 0.9),
]
)
docs = vector.search_by_vector([0.1], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"])
assert len(docs) == 1
assert vector.search_by_full_text("query") == []
# 8. delete commits session
def test_delete_drops_table(relyt_module, monkeypatch):
vector = relyt_module.RelytVector.__new__(relyt_module.RelytVector)
vector._collection_name = "collection_1"
vector.client = MagicMock()
vector.embedding_dimension = 3
session = _FakeSession()
_patch_both(monkeypatch, relyt_module, session)
vector.delete()
session.execute.assert_called_once()
def test_relyt_factory_existing_and_generated_collection(relyt_module, monkeypatch):
factory = relyt_module.RelytVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(relyt_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(relyt_module.dify_config, "RELYT_HOST", "localhost")
monkeypatch.setattr(relyt_module.dify_config, "RELYT_PORT", 5432)
monkeypatch.setattr(relyt_module.dify_config, "RELYT_USER", "postgres")
monkeypatch.setattr(relyt_module.dify_config, "RELYT_PASSWORD", "secret")
monkeypatch.setattr(relyt_module.dify_config, "RELYT_DATABASE", "relyt")
with patch.object(relyt_module, "RelytVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,316 +0,0 @@
import importlib
import json
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from pydantic import ValidationError
from core.rag.models.document import Document
def _build_fake_tablestore_module():
tablestore = types.ModuleType("tablestore")
class _BatchGetRowRequest:
def __init__(self):
self.items = []
def add(self, item):
self.items.append(item)
class _TableInBatchGetRowItem:
def __init__(self, table_name, rows_to_get, columns_to_get, _unused, _ver):
self.table_name = table_name
self.rows_to_get = rows_to_get
self.columns_to_get = columns_to_get
class _Row:
def __init__(self, primary_key, attribute_columns=None):
self.primary_key = primary_key
self.attribute_columns = attribute_columns or []
class _Client:
def __init__(self, *_args):
self.list_table = MagicMock(return_value=[])
self.create_table = MagicMock()
self.list_search_index = MagicMock(return_value=[])
self.create_search_index = MagicMock()
self.delete_search_index = MagicMock()
self.delete_table = MagicMock()
self.put_row = MagicMock()
self.delete_row = MagicMock()
self.get_row = MagicMock(return_value=(None, None, None))
self.batch_get_row = MagicMock()
self.search = MagicMock()
tablestore.OTSClient = _Client
tablestore.BatchGetRowRequest = _BatchGetRowRequest
tablestore.TableInBatchGetRowItem = _TableInBatchGetRowItem
tablestore.Row = _Row
tablestore.TableMeta = lambda name, schema: ("table_meta", name, schema)
tablestore.TableOptions = lambda: ("table_options",)
tablestore.CapacityUnit = lambda read, write: ("capacity", read, write)
tablestore.ReservedThroughput = lambda cap: ("reserved", cap)
tablestore.FieldSchema = lambda *args, **kwargs: ("field", args, kwargs)
tablestore.VectorOptions = lambda **kwargs: ("vector_options", kwargs)
tablestore.SearchIndexMeta = lambda field_schemas: ("search_index_meta", field_schemas)
tablestore.SearchQuery = lambda query, **kwargs: SimpleNamespace(query=query, **kwargs)
tablestore.TermQuery = lambda key, value: ("term_query", key, value)
tablestore.ColumnsToGet = lambda **kwargs: ("columns_to_get", kwargs)
tablestore.KnnVectorQuery = lambda **kwargs: SimpleNamespace(**kwargs)
tablestore.TermsQuery = lambda key, values: ("terms_query", key, values)
tablestore.Sort = lambda **kwargs: ("sort", kwargs)
tablestore.ScoreSort = lambda **kwargs: ("score_sort", kwargs)
tablestore.BoolQuery = lambda **kwargs: SimpleNamespace(**kwargs)
tablestore.MatchQuery = lambda **kwargs: ("match_query", kwargs)
tablestore.FieldType = SimpleNamespace(TEXT="TEXT", VECTOR="VECTOR", KEYWORD="KEYWORD")
tablestore.AnalyzerType = SimpleNamespace(MAXWORD="MAXWORD")
tablestore.VectorDataType = SimpleNamespace(VD_FLOAT_32="VD_FLOAT_32")
tablestore.VectorMetricType = SimpleNamespace(VM_COSINE="VM_COSINE")
tablestore.ColumnReturnType = SimpleNamespace(SPECIFIED="SPECIFIED", ALL_FROM_INDEX="ALL_FROM_INDEX")
tablestore.SortOrder = SimpleNamespace(DESC="DESC")
return tablestore
@pytest.fixture
def tablestore_module(monkeypatch):
fake_module = _build_fake_tablestore_module()
monkeypatch.setitem(sys.modules, "tablestore", fake_module)
import core.rag.datasource.vdb.tablestore.tablestore_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"access_key_id": "ak",
"access_key_secret": "sk",
"instance_name": "instance",
"endpoint": "endpoint",
"normalize_full_text_bm25_score": False,
}
values.update(overrides)
return module.TableStoreConfig.model_validate(values)
@pytest.mark.parametrize(
("field", "value", "message"),
[
("access_key_id", "", "config ACCESS_KEY_ID is required"),
("access_key_secret", "", "config ACCESS_KEY_SECRET is required"),
("instance_name", "", "config INSTANCE_NAME is required"),
("endpoint", "", "config ENDPOINT is required"),
],
)
def test_tablestore_config_validation(tablestore_module, field, value, message):
values = _config(tablestore_module).model_dump()
values[field] = value
with pytest.raises(ValidationError, match=message):
tablestore_module.TableStoreConfig.model_validate(values)
def test_init_and_basic_delegation(tablestore_module):
vector = tablestore_module.TableStoreVector("collection_1", _config(tablestore_module))
assert vector.get_type() == tablestore_module.VectorType.TABLESTORE
assert vector._table_name == "collection_1"
assert vector._index_name == "collection_1_idx"
vector._create_collection = MagicMock()
vector.add_texts = MagicMock()
docs = [Document(page_content="hello", metadata={"doc_id": "d-1"})]
vector.create(docs, [[0.1, 0.2]])
vector._create_collection.assert_called_once_with(2)
vector.add_texts.assert_called_once_with(documents=docs, embeddings=[[0.1, 0.2]])
vector.create_collection([[0.1, 0.2]])
assert vector._create_collection.call_count == 2
def test_get_by_ids_text_exists_delete_and_wrappers(tablestore_module):
vector = tablestore_module.TableStoreVector("collection_1", _config(tablestore_module))
# get_by_ids
ok_item = SimpleNamespace(
is_ok=True,
row=SimpleNamespace(
attribute_columns=[("metadata", json.dumps({"doc_id": "1"}), None), ("page_content", "text-1", None)]
),
)
fail_item = SimpleNamespace(is_ok=False, row=None)
batch_resp = SimpleNamespace(get_result_by_table=lambda _table: [ok_item, fail_item])
vector._tablestore_client.batch_get_row.return_value = batch_resp
docs = vector.get_by_ids(["id-1"])
assert len(docs) == 1
assert docs[0].page_content == "text-1"
# text_exists
vector._tablestore_client.get_row.return_value = (None, object(), None)
assert vector.text_exists("id-1") is True
vector._tablestore_client.get_row.return_value = (None, None, None)
assert vector.text_exists("id-1") is False
# delete wrappers
vector._delete_row = MagicMock()
vector.delete_by_ids([])
vector._delete_row.assert_not_called()
vector.delete_by_ids(["id-1", "id-2"])
assert vector._delete_row.call_count == 2
vector._search_by_metadata = MagicMock(return_value=["id-a"])
assert vector.get_ids_by_metadata_field("document_id", "doc-1") == ["id-a"]
vector.delete_by_ids = MagicMock()
vector.delete_by_metadata_field("document_id", "doc-1")
vector.delete_by_ids.assert_called_once_with(["id-a"])
vector._search_by_vector = MagicMock(return_value=["vec-doc"])
vector._search_by_full_text = MagicMock(return_value=["fts-doc"])
assert vector.search_by_vector([0.1], top_k=2, score_threshold=0.5, document_ids_filter=["d-1"]) == ["vec-doc"]
assert vector.search_by_full_text("query", top_k=2, score_threshold=0.3, document_ids_filter=["d-1"]) == ["fts-doc"]
vector._delete_table_if_exist = MagicMock()
vector.delete()
vector._delete_table_if_exist.assert_called_once()
def test_create_collection_and_table_index_lifecycle(tablestore_module, monkeypatch):
vector = tablestore_module.TableStoreVector("collection_1", _config(tablestore_module))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(tablestore_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(tablestore_module.redis_client, "set", MagicMock())
monkeypatch.setattr(tablestore_module.redis_client, "get", MagicMock(return_value=1))
vector._create_table_if_not_exist = MagicMock()
vector._create_search_index_if_not_exist = MagicMock()
vector._create_collection(3)
vector._create_table_if_not_exist.assert_not_called()
monkeypatch.setattr(tablestore_module.redis_client, "get", MagicMock(return_value=None))
vector._create_collection(3)
vector._create_table_if_not_exist.assert_called_once()
vector._create_search_index_if_not_exist.assert_called_once_with(3)
tablestore_module.redis_client.set.assert_called_once()
vector = tablestore_module.TableStoreVector("collection_2", _config(tablestore_module))
vector._tablestore_client.list_table.return_value = ["collection_2"]
assert vector._create_table_if_not_exist() is None
vector._tablestore_client.list_table.return_value = []
vector._create_table_if_not_exist()
vector._tablestore_client.create_table.assert_called_once()
vector._tablestore_client.list_search_index.return_value = [("collection_2", "collection_2_idx")]
assert vector._create_search_index_if_not_exist(3) is None
vector._tablestore_client.list_search_index.return_value = []
vector._create_search_index_if_not_exist(3)
vector._tablestore_client.create_search_index.assert_called_once()
vector._tablestore_client.list_search_index.return_value = [("collection_2", "idx_a"), ("collection_2", "idx_b")]
vector._delete_table_if_exist()
assert vector._tablestore_client.delete_search_index.call_count == 2
vector._tablestore_client.delete_table.assert_called_once_with("collection_2")
vector._delete_search_index()
vector._tablestore_client.delete_search_index.assert_called_with("collection_2", "collection_2_idx")
def test_write_row_and_search_helpers(tablestore_module):
vector = tablestore_module.TableStoreVector("collection_1", _config(tablestore_module))
vector._write_row(
"id-1",
{
"page_content": "hello",
"vector": [0.1, 0.2],
"metadata": {"doc_id": "d-1", "document_id": "doc-1"},
},
)
put_row_call = vector._tablestore_client.put_row.call_args
assert put_row_call.args[0] == "collection_1"
attrs = put_row_call.args[1].attribute_columns
assert any(item[0] == "metadata_tags" for item in attrs)
vector._delete_row("id-1")
vector._tablestore_client.delete_row.assert_called_once()
# metadata search pagination
first_page = SimpleNamespace(rows=[[(("id", "row-1"),)]], next_token=b"next")
second_page = SimpleNamespace(rows=[[(("id", "row-2"),)]], next_token=b"")
vector._tablestore_client.search.side_effect = [first_page, second_page]
ids = vector._search_by_metadata("document_id", "doc-1")
assert ids == ["row-1", "row-2"]
vector._tablestore_client.search.side_effect = None
# vector search
hit1 = SimpleNamespace(
score=0.9,
row=(
None,
[("page_content", "doc-a"), ("metadata", json.dumps({"doc_id": "1"})), ("vector", json.dumps([0.1]))],
),
)
hit2 = SimpleNamespace(
score=0.2,
row=(
None,
[("page_content", "doc-b"), ("metadata", json.dumps({"doc_id": "2"})), ("vector", json.dumps([0.2]))],
),
)
vector._tablestore_client.search.return_value = SimpleNamespace(search_hits=[hit1, hit2])
docs = vector._search_by_vector([0.1], document_ids_filter=["document_id=doc-1"], top_k=2, score_threshold=0.5)
assert len(docs) == 1
assert docs[0].metadata["score"] == pytest.approx(0.9)
assert tablestore_module.TableStoreVector._normalize_score_exp_decay(0) == pytest.approx(0.0)
assert tablestore_module.TableStoreVector._normalize_score_exp_decay(100) <= 1.0
# full text search with and without normalized score filter
vector._normalize_full_text_bm25_score = True
hit3 = SimpleNamespace(
score=10.0, row=(None, [("page_content", "doc-c"), ("metadata", json.dumps({"doc_id": "3"}))])
)
hit4 = SimpleNamespace(
score=0.1, row=(None, [("page_content", "doc-d"), ("metadata", json.dumps({"doc_id": "4"}))])
)
vector._tablestore_client.search.return_value = SimpleNamespace(search_hits=[hit3, hit4])
docs = vector._search_by_full_text("query", document_ids_filter=["document_id=doc-1"], top_k=2, score_threshold=0.2)
assert len(docs) == 1
assert "score" in docs[0].metadata
vector._normalize_full_text_bm25_score = False
vector._tablestore_client.search.return_value = SimpleNamespace(search_hits=[hit3])
docs = vector._search_by_full_text("query", document_ids_filter=None, top_k=2, score_threshold=0.0)
assert len(docs) == 1
assert "score" not in docs[0].metadata
def test_tablestore_factory_uses_existing_or_generated_collection(tablestore_module, monkeypatch):
factory = tablestore_module.TableStoreVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(tablestore_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(tablestore_module.dify_config, "TABLESTORE_ENDPOINT", "endpoint")
monkeypatch.setattr(tablestore_module.dify_config, "TABLESTORE_INSTANCE_NAME", "instance")
monkeypatch.setattr(tablestore_module.dify_config, "TABLESTORE_ACCESS_KEY_ID", "ak")
monkeypatch.setattr(tablestore_module.dify_config, "TABLESTORE_ACCESS_KEY_SECRET", "sk")
monkeypatch.setattr(tablestore_module.dify_config, "TABLESTORE_NORMALIZE_FULLTEXT_BM25_SCORE", True)
with patch.object(tablestore_module, "TableStoreVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "EXISTING_COLLECTION"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "AUTO_COLLECTION"
assert dataset_without_index.index_struct is not None
@@ -1,309 +0,0 @@
import importlib
import sys
import types
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
from core.rag.models.document import Document
def _build_fake_tencent_modules():
tcvdb_text = types.ModuleType("tcvdb_text")
tcvdb_text_encoder = types.ModuleType("tcvdb_text.encoder")
tcvectordb = types.ModuleType("tcvectordb")
tcvectordb_model = types.ModuleType("tcvectordb.model")
tcvectordb_document = types.ModuleType("tcvectordb.model.document")
tcvectordb_index = types.ModuleType("tcvectordb.model.index")
tcvectordb_enum = types.ModuleType("tcvectordb.model.enum")
class _BM25Encoder:
def encode_texts(self, text):
return {"encoded_text": text}
def encode_queries(self, query):
return {"encoded_query": query}
@classmethod
def default(cls, _lang):
return cls()
class VectorDBError(Exception):
def __init__(self, message):
super().__init__(message)
self.message = message
class RPCVectorDBClient:
def __init__(self, **kwargs):
self.kwargs = kwargs
self.create_database_if_not_exists = MagicMock()
self.exists_collection = MagicMock(return_value=False)
self.describe_collection = MagicMock(return_value=SimpleNamespace(indexes=[]))
self.create_collection = MagicMock()
self.upsert = MagicMock()
self.query = MagicMock(return_value=[])
self.delete = MagicMock()
self.search = MagicMock(return_value=[])
self.hybrid_search = MagicMock(return_value=[])
self.drop_collection = MagicMock()
class _Document:
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
class _HNSWSearchParams:
def __init__(self, ef):
self.ef = ef
class _AnnSearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
class _KeywordSearch:
def __init__(self, **kwargs):
self.kwargs = kwargs
class _WeightedRerank:
def __init__(self, **kwargs):
self.kwargs = kwargs
class _Filter:
@staticmethod
def in_(field, values):
return ("in", field, values)
def __init__(self, condition):
self.condition = condition
_Filter.In = staticmethod(_Filter.in_)
class _HNSWParams:
def __init__(self, **kwargs):
self.kwargs = kwargs
class _FilterIndex:
def __init__(self, *args):
self.args = args
class _VectorIndex:
def __init__(self, *args):
self.args = args
class _SparseIndex:
def __init__(self, **kwargs):
self.kwargs = kwargs
tcvectordb_enum.IndexType = SimpleNamespace(
__members__={"HNSW": "HNSW", "PRIMARY_KEY": "PRIMARY_KEY", "FILTER": "FILTER", "SPARSE_INVERTED": "SPARSE"},
PRIMARY_KEY="PRIMARY_KEY",
FILTER="FILTER",
SPARSE_INVERTED="SPARSE",
)
tcvectordb_enum.MetricType = SimpleNamespace(__members__={"IP": "IP"}, IP="IP")
tcvectordb_enum.FieldType = SimpleNamespace(String="String", Json="Json", SparseVector="SparseVector")
tcvectordb_document.Document = _Document
tcvectordb_document.HNSWSearchParams = _HNSWSearchParams
tcvectordb_document.AnnSearch = _AnnSearch
tcvectordb_document.Filter = _Filter
tcvectordb_document.KeywordSearch = _KeywordSearch
tcvectordb_document.WeightedRerank = _WeightedRerank
tcvectordb_index.HNSWParams = _HNSWParams
tcvectordb_index.FilterIndex = _FilterIndex
tcvectordb_index.VectorIndex = _VectorIndex
tcvectordb_index.SparseIndex = _SparseIndex
tcvdb_text_encoder.BM25Encoder = _BM25Encoder
tcvectordb_model.document = tcvectordb_document
tcvectordb_model.enum = tcvectordb_enum
tcvectordb_model.index = tcvectordb_index
tcvectordb.RPCVectorDBClient = RPCVectorDBClient
tcvectordb.VectorDBException = VectorDBError
return {
"tcvdb_text": tcvdb_text,
"tcvdb_text.encoder": tcvdb_text_encoder,
"tcvectordb": tcvectordb,
"tcvectordb.model": tcvectordb_model,
"tcvectordb.model.document": tcvectordb_document,
"tcvectordb.model.index": tcvectordb_index,
"tcvectordb.model.enum": tcvectordb_enum,
}
@pytest.fixture
def tencent_module(monkeypatch):
for name, module in _build_fake_tencent_modules().items():
monkeypatch.setitem(sys.modules, name, module)
import core.rag.datasource.vdb.tencent.tencent_vector as module
return importlib.reload(module)
def _config(module, **overrides):
values = {
"url": "http://vdb.local",
"api_key": "api-key",
"timeout": 30,
"username": "user",
"database": "db",
"index_type": "HNSW",
"metric_type": "IP",
"shard": 1,
"replicas": 2,
"max_upsert_batch_size": 2,
"enable_hybrid_search": False,
}
values.update(overrides)
return module.TencentConfig.model_validate(values)
def test_config_and_init_paths(tencent_module):
config = _config(tencent_module)
assert config.to_tencent_params()["url"] == "http://vdb.local"
vector = tencent_module.TencentVector("collection_1", config)
assert vector.get_type() == tencent_module.VectorType.TENCENT
assert vector._client.kwargs["key"] == "api-key"
vector._client.exists_collection.return_value = True
vector._client.describe_collection.return_value = SimpleNamespace(
indexes=[SimpleNamespace(name="vector", dimension=768), SimpleNamespace(name="sparse_vector", dimension=0)]
)
vector._client_config.enable_hybrid_search = True
vector._load_collection()
assert vector._enable_hybrid_search is True
assert vector._dimension == 768
vector._client.describe_collection.return_value = SimpleNamespace(
indexes=[SimpleNamespace(name="vector", dimension=512)]
)
vector._load_collection()
assert vector._enable_hybrid_search is False
def test_create_collection_branches(tencent_module, monkeypatch):
vector = tencent_module.TencentVector("collection_1", _config(tencent_module))
lock = MagicMock()
lock.__enter__.return_value = None
lock.__exit__.return_value = None
monkeypatch.setattr(tencent_module.redis_client, "lock", MagicMock(return_value=lock))
monkeypatch.setattr(tencent_module.redis_client, "set", MagicMock())
monkeypatch.setattr(tencent_module.redis_client, "get", MagicMock(return_value=1))
vector._create_collection(3)
vector._client.create_collection.assert_not_called()
monkeypatch.setattr(tencent_module.redis_client, "get", MagicMock(return_value=None))
vector._client.exists_collection.return_value = True
vector._create_collection(3)
vector._client.create_collection.assert_not_called()
vector._client.exists_collection.return_value = False
vector._client_config.index_type = "UNKNOWN"
with pytest.raises(ValueError, match="unsupported index_type"):
vector._create_collection(3)
vector._client_config.index_type = "HNSW"
vector._client_config.metric_type = "UNKNOWN"
with pytest.raises(ValueError, match="unsupported metric_type"):
vector._create_collection(3)
vector._client_config.metric_type = "IP"
vector._client.create_collection.side_effect = [
tencent_module.VectorDBException("fieldType:json unsupported"),
None,
]
vector._enable_hybrid_search = True
vector._create_collection(3)
assert vector._client.create_collection.call_count == 2
tencent_module.redis_client.set.assert_called_once()
vector._client.create_collection.side_effect = None
def test_create_add_delete_and_search_behaviour(tencent_module):
vector = tencent_module.TencentVector("collection_1", _config(tencent_module, enable_hybrid_search=True))
vector._create_collection = MagicMock()
docs = [
Document(page_content="text-a", metadata={"doc_id": "a", "document_id": "doc-a"}),
Document(page_content="text-b", metadata={"doc_id": "b", "document_id": "doc-b"}),
Document(page_content="text-c", metadata={"doc_id": "c", "document_id": "doc-c"}),
]
embeddings = [[0.1], [0.2], [0.3]]
vector.create(docs, embeddings)
vector._create_collection.assert_called_once_with(1)
vector._client.upsert.reset_mock()
vector.add_texts(docs, embeddings)
assert vector._client.upsert.call_count == 2
first_docs = vector._client.upsert.call_args_list[0].kwargs["documents"]
assert "sparse_vector" in first_docs[0].__dict__
vector._client.query.return_value = [{"id": "a"}]
assert vector.text_exists("a") is True
vector._client.query.return_value = []
assert vector.text_exists("a") is False
vector.delete_by_ids([])
vector._client.delete.assert_not_called()
vector.delete_by_ids(["a", "b", "c"])
assert vector._client.delete.call_count == 2
vector.delete_by_metadata_field("document_id", "doc-a")
assert vector._client.delete.call_count >= 3
vector._client.search.return_value = [[{"metadata": {"doc_id": "1"}, "text": "vec-doc", "score": 0.9}]]
vec_docs = vector.search_by_vector([0.1], top_k=2, score_threshold=0.5, document_ids_filter=["doc-a"])
assert len(vec_docs) == 1
assert vec_docs[0].metadata["score"] == pytest.approx(0.9)
vector._enable_hybrid_search = False
assert vector.search_by_full_text("query") == []
vector._enable_hybrid_search = True
vector._client.hybrid_search.return_value = [[{"metadata": {"doc_id": "2"}, "text": "fts-doc", "score": 0.8}]]
fts_docs = vector.search_by_full_text("query", top_k=2, score_threshold=0.5, document_ids_filter=["doc-a"])
assert len(fts_docs) == 1
# _get_search_res handles old string metadata format
compat_docs = vector._get_search_res([[{"metadata": '{"doc_id": "3"}', "text": "compat", "score": 0.2}]], 0.5)
assert len(compat_docs) == 1
assert compat_docs[0].metadata["score"] == pytest.approx(0.8)
vector._has_collection = MagicMock(return_value=True)
vector.delete()
vector._client.drop_collection.assert_called_once()
def test_tencent_factory_existing_and_generated_collection(tencent_module, monkeypatch):
factory = tencent_module.TencentVectorFactory()
dataset_with_index = SimpleNamespace(
id="dataset-1",
index_struct_dict={"vector_store": {"class_prefix": "EXISTING_COLLECTION"}},
index_struct=None,
)
dataset_without_index = SimpleNamespace(id="dataset-2", index_struct_dict=None, index_struct=None)
monkeypatch.setattr(tencent_module.Dataset, "gen_collection_name_by_id", lambda _id: "AUTO_COLLECTION")
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_URL", "http://vdb.local")
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_API_KEY", "api-key")
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_TIMEOUT", 30)
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_USERNAME", "user")
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_DATABASE", "db")
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_SHARD", 1)
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_REPLICAS", 2)
monkeypatch.setattr(tencent_module.dify_config, "TENCENT_VECTOR_DB_ENABLE_HYBRID_SEARCH", True)
with patch.object(tencent_module, "TencentVector", return_value="vector") as vector_cls:
result_1 = factory.init_vector(dataset_with_index, attributes=[], embeddings=MagicMock())
result_2 = factory.init_vector(dataset_without_index, attributes=[], embeddings=MagicMock())
assert result_1 == "vector"
assert result_2 == "vector"
assert vector_cls.call_args_list[0].kwargs["collection_name"] == "existing_collection"
assert vector_cls.call_args_list[1].kwargs["collection_name"] == "auto_collection"
assert dataset_without_index.index_struct is not None
@@ -21,6 +21,9 @@ def _register_fake_factory_module(monkeypatch, module_path: str, class_name: str
def vector_factory_module():
import importlib
from core.rag.datasource.vdb import vector_backend_registry as reg
reg.clear_vector_factory_cache()
import core.rag.datasource.vdb.vector_factory as module
return importlib.reload(module)
@@ -41,61 +44,62 @@ def test_gen_index_struct_dict(vector_factory_module):
@pytest.mark.parametrize(
("vector_type", "module_path", "class_name"),
[
("CHROMA", "core.rag.datasource.vdb.chroma.chroma_vector", "ChromaVectorFactory"),
("MILVUS", "core.rag.datasource.vdb.milvus.milvus_vector", "MilvusVectorFactory"),
("CHROMA", "dify_vdb_chroma.chroma_vector", "ChromaVectorFactory"),
("MILVUS", "dify_vdb_milvus.milvus_vector", "MilvusVectorFactory"),
(
"ALIBABACLOUD_MYSQL",
"core.rag.datasource.vdb.alibabacloud_mysql.alibabacloud_mysql_vector",
"dify_vdb_alibabacloud_mysql.alibabacloud_mysql_vector",
"AlibabaCloudMySQLVectorFactory",
),
("MYSCALE", "core.rag.datasource.vdb.myscale.myscale_vector", "MyScaleVectorFactory"),
("PGVECTOR", "core.rag.datasource.vdb.pgvector.pgvector", "PGVectorFactory"),
("VASTBASE", "core.rag.datasource.vdb.pyvastbase.vastbase_vector", "VastbaseVectorFactory"),
("PGVECTO_RS", "core.rag.datasource.vdb.pgvecto_rs.pgvecto_rs", "PGVectoRSFactory"),
("QDRANT", "core.rag.datasource.vdb.qdrant.qdrant_vector", "QdrantVectorFactory"),
("RELYT", "core.rag.datasource.vdb.relyt.relyt_vector", "RelytVectorFactory"),
("MYSCALE", "dify_vdb_myscale.myscale_vector", "MyScaleVectorFactory"),
("PGVECTOR", "dify_vdb_pgvector.pgvector", "PGVectorFactory"),
("VASTBASE", "dify_vdb_vastbase.vastbase_vector", "VastbaseVectorFactory"),
("PGVECTO_RS", "dify_vdb_pgvecto_rs.pgvecto_rs", "PGVectoRSFactory"),
("QDRANT", "dify_vdb_qdrant.qdrant_vector", "QdrantVectorFactory"),
("RELYT", "dify_vdb_relyt.relyt_vector", "RelytVectorFactory"),
(
"ELASTICSEARCH",
"core.rag.datasource.vdb.elasticsearch.elasticsearch_vector",
"dify_vdb_elasticsearch.elasticsearch_vector",
"ElasticSearchVectorFactory",
),
(
"ELASTICSEARCH_JA",
"core.rag.datasource.vdb.elasticsearch.elasticsearch_ja_vector",
"dify_vdb_elasticsearch.elasticsearch_ja_vector",
"ElasticSearchJaVectorFactory",
),
("TIDB_VECTOR", "core.rag.datasource.vdb.tidb_vector.tidb_vector", "TiDBVectorFactory"),
("WEAVIATE", "core.rag.datasource.vdb.weaviate.weaviate_vector", "WeaviateVectorFactory"),
("TENCENT", "core.rag.datasource.vdb.tencent.tencent_vector", "TencentVectorFactory"),
("ORACLE", "core.rag.datasource.vdb.oracle.oraclevector", "OracleVectorFactory"),
("TIDB_VECTOR", "dify_vdb_tidb_vector.tidb_vector", "TiDBVectorFactory"),
("WEAVIATE", "dify_vdb_weaviate.weaviate_vector", "WeaviateVectorFactory"),
("TENCENT", "dify_vdb_tencent.tencent_vector", "TencentVectorFactory"),
("ORACLE", "dify_vdb_oracle.oraclevector", "OracleVectorFactory"),
(
"OPENSEARCH",
"core.rag.datasource.vdb.opensearch.opensearch_vector",
"dify_vdb_opensearch.opensearch_vector",
"OpenSearchVectorFactory",
),
("ANALYTICDB", "core.rag.datasource.vdb.analyticdb.analyticdb_vector", "AnalyticdbVectorFactory"),
("COUCHBASE", "core.rag.datasource.vdb.couchbase.couchbase_vector", "CouchbaseVectorFactory"),
("BAIDU", "core.rag.datasource.vdb.baidu.baidu_vector", "BaiduVectorFactory"),
("VIKINGDB", "core.rag.datasource.vdb.vikingdb.vikingdb_vector", "VikingDBVectorFactory"),
("UPSTASH", "core.rag.datasource.vdb.upstash.upstash_vector", "UpstashVectorFactory"),
("ANALYTICDB", "dify_vdb_analyticdb.analyticdb_vector", "AnalyticdbVectorFactory"),
("COUCHBASE", "dify_vdb_couchbase.couchbase_vector", "CouchbaseVectorFactory"),
("BAIDU", "dify_vdb_baidu.baidu_vector", "BaiduVectorFactory"),
("VIKINGDB", "dify_vdb_vikingdb.vikingdb_vector", "VikingDBVectorFactory"),
("UPSTASH", "dify_vdb_upstash.upstash_vector", "UpstashVectorFactory"),
(
"TIDB_ON_QDRANT",
"core.rag.datasource.vdb.tidb_on_qdrant.tidb_on_qdrant_vector",
"dify_vdb_tidb_on_qdrant.tidb_on_qdrant_vector",
"TidbOnQdrantVectorFactory",
),
("LINDORM", "core.rag.datasource.vdb.lindorm.lindorm_vector", "LindormVectorStoreFactory"),
("OCEANBASE", "core.rag.datasource.vdb.oceanbase.oceanbase_vector", "OceanBaseVectorFactory"),
("SEEKDB", "core.rag.datasource.vdb.oceanbase.oceanbase_vector", "OceanBaseVectorFactory"),
("OPENGAUSS", "core.rag.datasource.vdb.opengauss.opengauss", "OpenGaussFactory"),
("TABLESTORE", "core.rag.datasource.vdb.tablestore.tablestore_vector", "TableStoreVectorFactory"),
("LINDORM", "dify_vdb_lindorm.lindorm_vector", "LindormVectorStoreFactory"),
("OCEANBASE", "dify_vdb_oceanbase.oceanbase_vector", "OceanBaseVectorFactory"),
("SEEKDB", "dify_vdb_oceanbase.oceanbase_vector", "OceanBaseVectorFactory"),
("OPENGAUSS", "dify_vdb_opengauss.opengauss", "OpenGaussFactory"),
("TABLESTORE", "dify_vdb_tablestore.tablestore_vector", "TableStoreVectorFactory"),
(
"HUAWEI_CLOUD",
"core.rag.datasource.vdb.huawei.huawei_cloud_vector",
"dify_vdb_huawei_cloud.huawei_cloud_vector",
"HuaweiCloudVectorFactory",
),
("MATRIXONE", "core.rag.datasource.vdb.matrixone.matrixone_vector", "MatrixoneVectorFactory"),
("CLICKZETTA", "core.rag.datasource.vdb.clickzetta.clickzetta_vector", "ClickzettaVectorFactory"),
("IRIS", "core.rag.datasource.vdb.iris.iris_vector", "IrisVectorFactory"),
("MATRIXONE", "dify_vdb_matrixone.matrixone_vector", "MatrixoneVectorFactory"),
("CLICKZETTA", "dify_vdb_clickzetta.clickzetta_vector", "ClickzettaVectorFactory"),
("IRIS", "dify_vdb_iris.iris_vector", "IrisVectorFactory"),
("HOLOGRES", "dify_vdb_hologres.hologres_vector", "HologresVectorFactory"),
],
)
def test_get_vector_factory_supported(vector_factory_module, monkeypatch, vector_type, module_path, class_name):
@@ -111,6 +115,34 @@ def test_get_vector_factory_unsupported(vector_factory_module):
vector_factory_module.Vector.get_vector_factory("unknown")
class _PluginChromaFactory:
"""Stub used only for entry-point override test."""
def test_get_vector_factory_entry_point_overrides_builtin(vector_factory_module, monkeypatch):
from importlib.metadata import EntryPoint
from core.rag.datasource.vdb import vector_backend_registry as reg
reg.clear_vector_factory_cache()
ep = EntryPoint(
name="chroma",
value=f"{__name__}:_PluginChromaFactory",
group="dify.vector_backends",
)
class _FakeGroups:
def select(self, *, group: str):
if group == "dify.vector_backends":
return (ep,)
return ()
monkeypatch.setattr(reg, "entry_points", lambda: _FakeGroups())
result_cls = vector_factory_module.Vector.get_vector_factory(vector_factory_module.VectorType.CHROMA)
assert result_cls is _PluginChromaFactory
def test_vector_init_uses_default_and_custom_attributes(vector_factory_module):
dataset = SimpleNamespace(id="dataset-1")
@@ -1,160 +0,0 @@
from unittest.mock import patch
import httpx
import pytest
from qdrant_client.http import models as rest
from qdrant_client.http.exceptions import UnexpectedResponse
from core.rag.datasource.vdb.tidb_on_qdrant.tidb_on_qdrant_vector import (
TidbOnQdrantConfig,
TidbOnQdrantVector,
)
class TestTidbOnQdrantVectorDeleteByIds:
"""Unit tests for TidbOnQdrantVector.delete_by_ids method."""
@pytest.fixture
def vector_instance(self):
"""Create a TidbOnQdrantVector instance for testing."""
config = TidbOnQdrantConfig(
endpoint="http://localhost:6333",
api_key="test_api_key",
)
with patch("core.rag.datasource.vdb.tidb_on_qdrant.tidb_on_qdrant_vector.qdrant_client.QdrantClient"):
vector = TidbOnQdrantVector(
collection_name="test_collection",
group_id="test_group",
config=config,
)
return vector
def test_delete_by_ids_with_multiple_ids(self, vector_instance):
"""Test batch deletion with multiple document IDs."""
ids = ["doc1", "doc2", "doc3"]
vector_instance.delete_by_ids(ids)
# Verify that delete was called once with MatchAny filter
vector_instance._client.delete.assert_called_once()
call_args = vector_instance._client.delete.call_args
# Check collection name
assert call_args[1]["collection_name"] == "test_collection"
# Verify filter uses MatchAny with all IDs
filter_selector = call_args[1]["points_selector"]
filter_obj = filter_selector.filter
assert len(filter_obj.must) == 1
field_condition = filter_obj.must[0]
assert field_condition.key == "metadata.doc_id"
assert isinstance(field_condition.match, rest.MatchAny)
assert set(field_condition.match.any) == {"doc1", "doc2", "doc3"}
def test_delete_by_ids_with_single_id(self, vector_instance):
"""Test deletion with a single document ID."""
ids = ["doc1"]
vector_instance.delete_by_ids(ids)
# Verify that delete was called once
vector_instance._client.delete.assert_called_once()
call_args = vector_instance._client.delete.call_args
# Verify filter uses MatchAny with single ID
filter_selector = call_args[1]["points_selector"]
filter_obj = filter_selector.filter
field_condition = filter_obj.must[0]
assert isinstance(field_condition.match, rest.MatchAny)
assert field_condition.match.any == ["doc1"]
def test_delete_by_ids_with_empty_list(self, vector_instance):
"""Test deletion with empty ID list returns early without API call."""
vector_instance.delete_by_ids([])
# Verify that delete was NOT called
vector_instance._client.delete.assert_not_called()
def test_delete_by_ids_with_404_error(self, vector_instance):
"""Test that 404 errors (collection not found) are handled gracefully."""
ids = ["doc1", "doc2"]
# Mock a 404 error
error = UnexpectedResponse(
status_code=404,
reason_phrase="Not Found",
content=b"Collection not found",
headers=httpx.Headers(),
)
vector_instance._client.delete.side_effect = error
# Should not raise an exception
vector_instance.delete_by_ids(ids)
# Verify delete was called
vector_instance._client.delete.assert_called_once()
def test_delete_by_ids_with_unexpected_error(self, vector_instance):
"""Test that non-404 errors are re-raised."""
ids = ["doc1", "doc2"]
# Mock a 500 error
error = UnexpectedResponse(
status_code=500,
reason_phrase="Internal Server Error",
content=b"Server error",
headers=httpx.Headers(),
)
vector_instance._client.delete.side_effect = error
# Should re-raise the exception
with pytest.raises(UnexpectedResponse) as exc_info:
vector_instance.delete_by_ids(ids)
assert exc_info.value.status_code == 500
def test_delete_by_ids_with_large_batch(self, vector_instance):
"""Test deletion with a large batch of IDs."""
# Create 1000 IDs
ids = [f"doc_{i}" for i in range(1000)]
vector_instance.delete_by_ids(ids)
# Verify single delete call with all IDs
vector_instance._client.delete.assert_called_once()
call_args = vector_instance._client.delete.call_args
filter_selector = call_args[1]["points_selector"]
filter_obj = filter_selector.filter
field_condition = filter_obj.must[0]
# Verify all 1000 IDs are in the batch
assert len(field_condition.match.any) == 1000
assert "doc_0" in field_condition.match.any
assert "doc_999" in field_condition.match.any
def test_delete_by_ids_filter_structure(self, vector_instance):
"""Test that the filter structure is correctly constructed."""
ids = ["doc1", "doc2"]
vector_instance.delete_by_ids(ids)
call_args = vector_instance._client.delete.call_args
filter_selector = call_args[1]["points_selector"]
filter_obj = filter_selector.filter
# Verify Filter structure
assert isinstance(filter_obj, rest.Filter)
assert filter_obj.must is not None
assert len(filter_obj.must) == 1
# Verify FieldCondition structure
field_condition = filter_obj.must[0]
assert isinstance(field_condition, rest.FieldCondition)
assert field_condition.key == "metadata.doc_id"
# Verify MatchAny structure
assert isinstance(field_condition.match, rest.MatchAny)
assert field_condition.match.any == ids

Some files were not shown because too many files have changed in this diff Show More