244 lines
7.9 KiB
Python
244 lines
7.9 KiB
Python
"""M6 Qdrant 向量存储模块测试。
|
|
|
|
使用 qdrant-client 内存模式(:memory:),不依赖 Docker。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
|
|
from vectorstore import (
|
|
SearchFilter,
|
|
SearchResult,
|
|
VectorStore,
|
|
make_qdrant_client,
|
|
)
|
|
from vectorstore.client import _build_filter
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# fixtures
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
@pytest.fixture
|
|
def store() -> VectorStore:
|
|
c = make_qdrant_client(memory=True)
|
|
s = VectorStore(c, collection_name="test_m6", vector_dim=4)
|
|
s.init_collection()
|
|
yield s
|
|
s.close()
|
|
|
|
|
|
def _point(id_: str, vector: list[float], **payload: object) -> dict:
|
|
return {"id": id_, "vector": vector, "payload": dict(payload)}
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# Collection 管理
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_init_collection_creates(store: VectorStore) -> None:
|
|
info = store.info()
|
|
assert info.exists is True
|
|
assert info.name == "test_m6"
|
|
|
|
|
|
def test_init_collection_idempotent(store: VectorStore) -> None:
|
|
"""再次 init 不应报错,count 不变。"""
|
|
store.upsert([_point("a", [1.0, 0, 0, 0])])
|
|
store.init_collection() # 不应重建
|
|
assert store.count() == 1
|
|
|
|
|
|
def test_init_collection_recreate_clears(store: VectorStore) -> None:
|
|
store.upsert([_point("a", [1.0, 0, 0, 0])])
|
|
store.init_collection(recreate=True)
|
|
assert store.count() == 0
|
|
|
|
|
|
def test_delete_collection(store: VectorStore) -> None:
|
|
store.delete_collection()
|
|
assert store.info().exists is False
|
|
# 再次 init 应恢复
|
|
store.init_collection()
|
|
assert store.info().exists is True
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# upsert + count
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_upsert_and_count(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], title="Article A"),
|
|
_point("b", [0, 1, 0, 0], title="Article B"),
|
|
])
|
|
assert store.count() == 2
|
|
|
|
|
|
def test_upsert_idempotent(store: VectorStore) -> None:
|
|
"""同 url_hash 再次 upsert 不应增加 count,数据被覆盖。"""
|
|
store.upsert([_point("a", [1, 0, 0, 0], title="Old")])
|
|
store.upsert([_point("a", [0, 0, 0, 1], title="New")])
|
|
assert store.count() == 1
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# query - 语义检索
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_query_returns_score_desc(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1.0, 0, 0, 0], title="A"),
|
|
_point("b", [0.0, 1.0, 0, 0], title="B"),
|
|
_point("c", [0.0, 0, 1.0, 0], title="C"),
|
|
])
|
|
results = store.query(query_vector=[0.9, 0.1, 0, 0], top_k=2)
|
|
assert len(results) == 2
|
|
assert results[0].url_hash == "a"
|
|
# score 应递减
|
|
assert results[0].score >= results[1].score
|
|
|
|
|
|
def test_query_score_threshold(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], title="A"),
|
|
_point("b", [0, 1, 0, 0], title="B"),
|
|
])
|
|
# 只有 a 会匹配
|
|
results = store.query(query_vector=[1, 0, 0, 0], top_k=10, score_threshold=0.9)
|
|
assert len(results) == 1
|
|
assert results[0].url_hash == "a"
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# query - 结构化过滤
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_query_filter_by_source_id(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a1", [1, 0, 0, 0], source_id="cls", title="CLS article"),
|
|
_point("a2", [0.9, 0.1, 0, 0], source_id="sina", title="Sina article"),
|
|
])
|
|
results = store.query(
|
|
query_vector=[1, 0, 0, 0],
|
|
filter=SearchFilter(source_id="cls"),
|
|
top_k=5,
|
|
)
|
|
assert len(results) == 1
|
|
assert results[0].source_id == "cls"
|
|
|
|
|
|
def test_query_filter_by_stock_codes(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], source_id="cls",
|
|
event={"stock_codes": ["300750"], "sentiment": "positive"}),
|
|
_point("b", [0.9, 0.1, 0, 0], source_id="sina",
|
|
event={"stock_codes": ["000001"], "sentiment": "neutral"}),
|
|
_point("c", [0.8, 0.2, 0, 0], source_id="sina",
|
|
event={"stock_codes": ["300750"], "sentiment": "negative"}),
|
|
])
|
|
results = store.query(
|
|
query_vector=[1, 0, 0, 0],
|
|
filter=SearchFilter(stock_codes=["300750"]),
|
|
top_k=5,
|
|
)
|
|
assert len(results) == 2
|
|
for r in results:
|
|
assert "300750" in (r.event or {}).get("stock_codes", [])
|
|
|
|
|
|
def test_query_filter_by_sentiment(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], source_id="cls",
|
|
event={"sentiment": "positive"}),
|
|
_point("b", [0, 1, 0, 0], source_id="cls",
|
|
event={"sentiment": "negative"}),
|
|
])
|
|
results = store.query(
|
|
query_vector=[1, 0, 0, 0],
|
|
filter=SearchFilter(sentiment="positive"),
|
|
top_k=5,
|
|
)
|
|
assert len(results) >= 1
|
|
assert all((r.event or {}).get("sentiment") == "positive" for r in results)
|
|
|
|
|
|
def test_query_filter_by_importance_min(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], source_id="cls",
|
|
event={"importance": 2}),
|
|
_point("b", [0, 1, 0, 0], source_id="cls",
|
|
event={"importance": 4}),
|
|
_point("c", [0, 0, 1, 0], source_id="cls",
|
|
event={"importance": 5}),
|
|
])
|
|
results = store.query(
|
|
query_vector=[0.5, 0.5, 0.5, 0],
|
|
filter=SearchFilter(importance_min=4),
|
|
top_k=5,
|
|
)
|
|
assert all((r.event or {}).get("importance", 0) >= 4 for r in results)
|
|
|
|
|
|
def test_query_filter_by_industry(store: VectorStore) -> None:
|
|
store.upsert([
|
|
_point("a", [1, 0, 0, 0], source_id="cls",
|
|
event={"industries": ["动力电池"]}),
|
|
_point("b", [0, 1, 0, 0], source_id="cls",
|
|
event={"industries": ["白酒"]}),
|
|
])
|
|
results = store.query(
|
|
query_vector=[1, 0, 0, 0],
|
|
filter=SearchFilter(industries=["动力电池"]),
|
|
top_k=5,
|
|
)
|
|
assert len(results) >= 1
|
|
for r in results:
|
|
assert "动力电池" in (r.event or {}).get("industries", [])
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# Filter 构建
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_build_filter_empty_returns_none() -> None:
|
|
assert _build_filter(SearchFilter()) is None
|
|
|
|
|
|
def test_build_filter_source_id() -> None:
|
|
f = _build_filter(SearchFilter(source_id="cls"))
|
|
assert f is not None and len(f.must) == 1 # type: ignore[arg-type]
|
|
|
|
|
|
def test_build_filter_date_range() -> None:
|
|
f = _build_filter(SearchFilter(publish_date_from="2026-06-01", publish_date_to="2026-06-30"))
|
|
assert f is not None
|
|
# range 应含 gte + lte
|
|
cond = f.must[0] # type: ignore[union-attr]
|
|
assert cond.key == "publish_time"
|
|
|
|
|
|
# --------------------------------------------------------------------------- #
|
|
# search_result 模型
|
|
# --------------------------------------------------------------------------- #
|
|
|
|
def test_search_result_short_summary() -> None:
|
|
r = SearchResult(
|
|
url_hash="abc",
|
|
score=0.95,
|
|
title="宁德时代签约 100GWh 协议",
|
|
url="https://x/1",
|
|
source_id="cls",
|
|
event={"stock_codes": ["300750"], "sentiment": "positive"},
|
|
)
|
|
s = r.short_summary()
|
|
assert "cls" in s and "0.9500" in s and "300750" in s
|
|
|
|
|
|
def test_search_result_handles_none_event() -> None:
|
|
r = SearchResult(
|
|
url_hash="abc", score=0.5, title="t", url="u", source_id="cls", event=None,
|
|
)
|
|
s = r.short_summary()
|
|
assert "-" in s # 无 stock_code
|