mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
autofix-ci[bot]
wangxiaolei
parent
c34f67495c
commit
ae898652b2
@@ -0,0 +1 @@
|
||||
"""Test suite root package (enables ``import tests.integration_tests...`` with ``pythonpath = .``)."""
|
||||
|
||||
@@ -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()
|
||||
-74
@@ -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
|
||||
-733
@@ -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)
|
||||
-384
@@ -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()
|
||||
-427
@@ -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
|
||||
-121
@@ -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
|
||||
-405
@@ -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")
|
||||
|
||||
|
||||
-160
@@ -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
Reference in New Issue
Block a user