test_device_retrieval.py 5.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154
  1. from __future__ import annotations
  2. from datetime import UTC, datetime
  3. import pytest
  4. def _row(**overrides):
  5. from app.core.knowledge.retrieval.device import DeviceSearchRow
  6. values = {
  7. "asset_uid": "11111111-1111-4111-8111-111111111111",
  8. "asset_type": "device",
  9. "name": "一号循环泵",
  10. "current_version": 3,
  11. "location": "动力车间",
  12. "organization": "设备动力部",
  13. "responsible_person": "张工",
  14. "source_codes": ("EQ-001", "ERP-PUMP-01"),
  15. "related_events": (
  16. ("fault", "轴承润滑故障", "FT-001"),
  17. ("alarm", "轴承温度持续升高", "AL-001"),
  18. ),
  19. "business_domain_uid": "22222222-2222-4222-8222-222222222222",
  20. "updated_at": datetime(2026, 7, 29, 8, 0, tzinfo=UTC),
  21. "rank": 0.95,
  22. }
  23. values.update(overrides)
  24. return DeviceSearchRow(**values)
  25. def test_device_query_is_trimmed_bounded_and_required():
  26. from app.core.knowledge.retrieval.device import normalize_device_query
  27. assert normalize_device_query(" EQ-001 ") == "EQ-001"
  28. with pytest.raises(ValueError, match="不能为空"):
  29. normalize_device_query(" ")
  30. with pytest.raises(ValueError, match="300"):
  31. normalize_device_query("x" * 301)
  32. @pytest.mark.parametrize(
  33. ("query", "expected_term"),
  34. [
  35. ("循环泵最近有哪些故障?", "循环泵"),
  36. ("动力车间有哪些设备?", "动力车间"),
  37. ("张工负责哪些设备?", "张工"),
  38. ("一号循环泵的责任人是谁?", "一号循环泵"),
  39. ("源 ID 为 EQ-001 的设备责任人是谁?", "EQ-001"),
  40. (
  41. "平台 UID:11111111-1111-4111-8111-111111111111",
  42. "11111111-1111-4111-8111-111111111111",
  43. ),
  44. ],
  45. )
  46. def test_device_query_terms_extract_searchable_entity_from_common_questions(
  47. query,
  48. expected_term,
  49. ):
  50. from app.core.knowledge.retrieval.device import device_query_terms
  51. terms = device_query_terms(query)
  52. assert terms[0] == expected_term
  53. assert normalize_device_query_for_assertion(query) in terms
  54. def normalize_device_query_for_assertion(query):
  55. from app.core.knowledge.retrieval.device import normalize_device_query
  56. return normalize_device_query(query)
  57. def test_device_evidence_is_stable_readable_and_does_not_expose_raw_payloads():
  58. from app.core.knowledge.retrieval.device import build_device_evidence
  59. evidence = build_device_evidence(_row())
  60. assert evidence.chunk_id == (
  61. "device:11111111-1111-4111-8111-111111111111:v3"
  62. )
  63. assert evidence.object_type == "DeviceAsset"
  64. assert evidence.object_version == 3
  65. assert evidence.point_keys == (
  66. "DeviceAsset/11111111-1111-4111-8111-111111111111/summary",
  67. )
  68. assert evidence.point_revisions == (3,)
  69. assert evidence.score == 0.95
  70. assert "设备名称:一号循环泵" in evidence.content
  71. assert "源 ID:EQ-001、ERP-PUMP-01" in evidence.content
  72. assert "位置:动力车间" in evidence.content
  73. assert "责任人:张工" in evidence.content
  74. assert "故障:轴承润滑故障(FT-001)" in evidence.content
  75. assert "告警:轴承温度持续升高(AL-001)" in evidence.content
  76. for forbidden in ("password", "permission_scope", "source_config"):
  77. assert forbidden not in evidence.content
  78. def test_device_evidence_is_deterministic_for_unsorted_duplicate_inputs():
  79. from app.core.knowledge.retrieval.device import build_device_evidence
  80. first = build_device_evidence(
  81. _row(
  82. source_codes=("ERP-PUMP-01", "EQ-001", "EQ-001"),
  83. related_events=(
  84. ("alarm", "轴承温度持续升高", "AL-001"),
  85. ("fault", "轴承润滑故障", "FT-001"),
  86. ("fault", "轴承润滑故障", "FT-001"),
  87. ),
  88. )
  89. )
  90. second = build_device_evidence(_row())
  91. assert first.content == second.content
  92. assert first.point_keys == second.point_keys
  93. def test_device_result_limit_is_bounded_before_repository_query():
  94. from app.core.knowledge.access import KnowledgeAccessContext
  95. from app.core.knowledge.retrieval.device import SqlDeviceKnowledgeRetriever
  96. class Repository:
  97. def __init__(self):
  98. self.calls = []
  99. def search(self, **kwargs):
  100. self.calls.append(kwargs)
  101. return [_row()]
  102. repository = Repository()
  103. retriever = SqlDeviceKnowledgeRetriever(repository)
  104. context = KnowledgeAccessContext(
  105. subject_id="viewer-1",
  106. roles=frozenset({"viewer"}),
  107. permissions=frozenset({"governance:read"}),
  108. business_domain_uids=frozenset(
  109. {"22222222-2222-4222-8222-222222222222"}
  110. ),
  111. correlation_id="correlation-1",
  112. )
  113. result = retriever.retrieve("循环泵", context, 999)
  114. assert len(result) == 1
  115. assert repository.calls == [
  116. {
  117. "query": "循环泵",
  118. "global_access": False,
  119. "business_domain_uids": (
  120. "22222222-2222-4222-8222-222222222222",
  121. ),
  122. "limit": 100,
  123. }
  124. ]