test_milvus_store.py 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116
  1. from app.embeddings import HashEmbeddingProvider
  2. from app.milvus_store import IndexedDocument, MilvusVectorStore
  3. class FakeSchema:
  4. def __init__(self) -> None:
  5. self.fields = []
  6. def add_field(self, *args, **kwargs) -> None:
  7. self.fields.append((args, kwargs))
  8. class FakeIndexParams:
  9. def __init__(self) -> None:
  10. self.indexes = []
  11. def add_index(self, **kwargs) -> None:
  12. self.indexes.append(kwargs)
  13. class FakeMilvusClient:
  14. last_instance = None
  15. def __init__(self, **kwargs) -> None:
  16. self.kwargs = kwargs
  17. self.created = None
  18. self.inserted = []
  19. self.search_args = None
  20. self.index_params = FakeIndexParams()
  21. FakeMilvusClient.last_instance = self
  22. @classmethod
  23. def create_schema(cls, **kwargs):
  24. return FakeSchema()
  25. def list_collections(self):
  26. return []
  27. def has_collection(self, **kwargs):
  28. return False
  29. def prepare_index_params(self):
  30. return self.index_params
  31. def create_collection(self, **kwargs):
  32. self.created = kwargs
  33. def load_collection(self, **kwargs):
  34. return None
  35. def insert(self, **kwargs):
  36. self.inserted = kwargs["data"]
  37. def flush(self, **kwargs):
  38. return None
  39. def search(self, **kwargs):
  40. self.search_args = kwargs
  41. return [
  42. [
  43. {
  44. "id": 1,
  45. "distance": 0.88,
  46. "entity": {
  47. "content": "耳机拆封且影响卫生安全时不适用无理由退货。",
  48. "source": "return_policy.md",
  49. "doc_type": "commerce_policy",
  50. "chunk_index": 1,
  51. "policy_type": "return_policy",
  52. "category": "all",
  53. "metadata": {"title": "商品退货政策"},
  54. },
  55. }
  56. ]
  57. ]
  58. def test_milvus_schema_insert_and_filtered_search(monkeypatch) -> None:
  59. monkeypatch.setattr("app.milvus_store.MilvusClient", FakeMilvusClient)
  60. embeddings = HashEmbeddingProvider(128)
  61. store = MilvusVectorStore(
  62. uri="http://localhost:19530",
  63. token="",
  64. collection_name="test_docs",
  65. embeddings=embeddings,
  66. )
  67. store.create_collection(recreate=True)
  68. client = FakeMilvusClient.last_instance
  69. assert client.created["collection_name"] == "test_docs"
  70. assert {item["index_type"] for item in client.index_params.indexes} == {
  71. "HNSW",
  72. "INVERTED",
  73. }
  74. inserted = store.insert_documents(
  75. [
  76. IndexedDocument(
  77. content="耳机拆封后不适用无理由退货",
  78. source="return_policy.md",
  79. doc_type="commerce_policy",
  80. chunk_index=1,
  81. policy_type="return_policy",
  82. category="all",
  83. metadata={},
  84. )
  85. ]
  86. )
  87. assert inserted == 1
  88. assert len(client.inserted[0]["embedding"]) == 128
  89. evidence = store.search("耳机拆封后能退货吗", top_k=4, policy_type="return_policy")
  90. assert client.search_args["filter"] == 'policy_type == "return_policy"'
  91. assert client.search_args["limit"] == 4
  92. assert evidence[0].source_type == "milvus"
  93. assert evidence[0].score == 0.88