test_knowledge_retrieval.py 4.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127
  1. from __future__ import annotations
  2. import os
  3. import uuid
  4. import pytest
  5. from sqlalchemy import create_engine, text
  6. from sqlalchemy.orm import Session
  7. from app.core.knowledge.access import KnowledgeAccessContext
  8. from app.core.knowledge.point_builder import build_knowledge_snapshot
  9. from app.core.knowledge.repository import SqlKnowledgeRepository
  10. from app.core.knowledge.retrieval.sql import SqlLexicalRetriever, SqlVectorRetriever
  11. from app.core.knowledge.sync import KnowledgeSyncService
  12. pytestmark = pytest.mark.integration
  13. class FixedEmbeddingProvider:
  14. profile_key = "qwen:retrieval-test:1024"
  15. def embed(self, texts: list[str]) -> list[list[float]]:
  16. return [[1.0] * 1024 for _text in texts]
  17. @pytest.fixture()
  18. def database_url():
  19. value = os.environ.get("TEST_DATABASE_URL")
  20. if not value:
  21. pytest.skip("TEST_DATABASE_URL is not configured")
  22. return value
  23. def test_sql_retrievers_filter_business_domain_before_recall(database_url):
  24. engine = create_engine(database_url)
  25. profile_id = str(uuid.uuid4())
  26. domain_a, domain_b = str(uuid.uuid4()), str(uuid.uuid4())
  27. source_a, source_b = str(uuid.uuid4()), str(uuid.uuid4())
  28. sources = ((source_a, domain_a, "customer_a"), (source_b, domain_b, "customer_b"))
  29. try:
  30. with engine.begin() as connection:
  31. connection.execute(
  32. text(
  33. """
  34. INSERT INTO public.knowledge_embedding_profiles (
  35. id, provider, model, dimension, status, config_hash, activated_at
  36. ) VALUES (
  37. CAST(:id AS uuid), 'qwen', 'retrieval-test', 1024, 'active',
  38. :config_hash, CURRENT_TIMESTAMP
  39. )
  40. """
  41. ),
  42. {"id": profile_id, "config_hash": "b" * 64},
  43. )
  44. for source_uid, domain_uid, name in sources:
  45. with Session(engine) as session, session.begin():
  46. embedder = FixedEmbeddingProvider()
  47. repository = SqlKnowledgeRepository(
  48. session,
  49. embedding_profile_id=profile_id,
  50. embedding_profile_key=embedder.profile_key,
  51. generation=1,
  52. workspace=f"dataops-global-{domain_uid}-g1",
  53. )
  54. snapshot = build_knowledge_snapshot(
  55. "DataFlow",
  56. {
  57. "uid": source_uid,
  58. "version": 1,
  59. "name": name,
  60. "purpose": "sync customer data",
  61. "business_domain_uid": domain_uid,
  62. },
  63. )
  64. KnowledgeSyncService(repository=repository, embedder=embedder).sync(
  65. snapshot
  66. )
  67. context = KnowledgeAccessContext(
  68. subject_id="user-a",
  69. roles=frozenset({"viewer"}),
  70. permissions=frozenset({"governance:read"}),
  71. business_domain_uids=frozenset({domain_a}),
  72. correlation_id="retrieval-test",
  73. )
  74. with Session(engine) as session:
  75. lexical = SqlLexicalRetriever(session).retrieve("customer_a", context, 20)
  76. vector = SqlVectorRetriever(session, FixedEmbeddingProvider()).retrieve(
  77. "customer", context, 20
  78. )
  79. assert lexical
  80. assert {item.object_uid for item in lexical} == {source_a}
  81. assert vector
  82. assert {item.business_domain_uid for item in vector} == {domain_a}
  83. assert all(item.point_revisions for item in vector)
  84. finally:
  85. with engine.begin() as connection:
  86. connection.execute(
  87. text(
  88. "DELETE FROM public.governance_documents "
  89. "WHERE object_uid = ANY(CAST(:uids AS uuid[]))"
  90. ),
  91. {"uids": [source_a, source_b]},
  92. )
  93. connection.execute(
  94. text(
  95. "DELETE FROM public.knowledge_points "
  96. "WHERE source_uid = ANY(CAST(:uids AS uuid[]))"
  97. ),
  98. {"uids": [source_a, source_b]},
  99. )
  100. connection.execute(
  101. text(
  102. "DELETE FROM public.knowledge_change_sets "
  103. "WHERE source_uid = ANY(CAST(:uids AS uuid[]))"
  104. ),
  105. {"uids": [source_a, source_b]},
  106. )
  107. connection.execute(
  108. text(
  109. "DELETE FROM public.knowledge_embedding_profiles "
  110. "WHERE id = CAST(:id AS uuid)"
  111. ),
  112. {"id": profile_id},
  113. )
  114. engine.dispose()