test_rag_retrieval.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256
  1. """
  2. Dify RAG Retrieval Test Script (Dual Mode)
  3. ===========================================
  4. Tests knowledge base retrieval across 5 query types.
  5. Modes:
  6. --mode local Local retriever (no Dify needed)
  7. --mode dify Dify API (requires API key)
  8. Usage:
  9. python test_rag_retrieval.py --mode local
  10. python test_rag_retrieval.py --mode dify --api-key <KEY>
  11. """
  12. import os
  13. import sys
  14. import json
  15. import argparse
  16. # (type, chinese_query, expected_file_substring, doc_hint)
  17. TEST_QUERIES = [
  18. # Type 1: Position queries -> should retrieve A1 (命盘24位置详解)
  19. ("position", "O位置代表什么含义", "命盘", "A1"),
  20. ("position", "P位置如何计算", "命盘", "A1"),
  21. ("position", "I和J位置代表什么能量", "命盘", "A1"),
  22. # Type 2: Combo pair queries -> should retrieve B2 (三角命盘组合对判读规则)
  23. ("combo", "M和N组合出现天医星怎么解读", "判读", "B2"),
  24. ("combo", "横向组合对和纵向组合对的区", "判读", "B2"),
  25. ("combo", "组合对分析的优先顺序是什么", "判读", "B2"),
  26. # Type 3: Dimension queries -> should retrieve C1 (数字能量分析维度手册)
  27. ("dimension", "天医星财富维度如何分析", "维度", "C1"),
  28. ("dimension", "延年星在事业维度的含义", "维度", "C1"),
  29. ("dimension", "五鬼星对健康维度的影响", "维度", "C1"),
  30. # Type 4: Zone/group queries -> should retrieve C2 (五区三组分析指南)
  31. ("zone", "父源区包含哪些位置", "五区", "C2"),
  32. ("zone", "左侧组年龄段和人生课题", "五区", "C2"),
  33. ("zone", "五区和三组的交叉分析", "五区", "C2"),
  34. # Type 5: Comprehensive
  35. ("comprehensive", "主性格6号人的感情事业特点", "维度", "C1"),
  36. ]
  37. def _dify_retrieve(query_text, api_key, base_url, dataset_id, top_k=6):
  38. """Dify API retrieval"""
  39. import requests
  40. headers = {
  41. "Authorization": f"Bearer {api_key}",
  42. "Content-Type": "application/json"
  43. }
  44. url = f"{base_url}/v1/datasets/{dataset_id}/retrieve"
  45. payload = {
  46. "query": query_text,
  47. "retrieval_model": {
  48. "search_method": "hybrid_search",
  49. "reranking_enable": True,
  50. "top_k": top_k,
  51. "score_threshold_enabled": False
  52. }
  53. }
  54. try:
  55. resp = requests.post(url, headers=headers, json=payload, timeout=30)
  56. if resp.status_code == 200:
  57. records = resp.json().get("records", [])
  58. return records, None
  59. else:
  60. return None, f"HTTP {resp.status_code} {resp.text[:200]}"
  61. except Exception as e:
  62. return None, str(e)
  63. def _dify_run_all(api_key, base_url, dataset_id):
  64. """Execute all Dify retrieval tests"""
  65. import requests
  66. print("=" * 70)
  67. print("Dify KB RAG Retrieval Test")
  68. print("=" * 70)
  69. print(f"Dataset: {dataset_id}")
  70. print(f"URL: {base_url}")
  71. # Show config
  72. try:
  73. resp = requests.get(f"{base_url}/v1/datasets/{dataset_id}",
  74. headers={"Authorization": f"Bearer {api_key}"}, timeout=15)
  75. if resp.status_code == 200:
  76. info = resp.json()
  77. rm = info.get("retrieval_model_dict", {})
  78. print(f"KB: {info.get('name', '?')}")
  79. print(f"Docs: {info.get('document_count')}")
  80. print(f"Config: method={rm.get('search_method')} top_k={rm.get('top_k')}")
  81. except Exception:
  82. pass
  83. print(f"Queries: {len(TEST_QUERIES)}")
  84. print()
  85. stats = {"pass": 0, "fail": 0, "error": 0, "total": len(TEST_QUERIES), "mode": "dify"}
  86. for qtype, query, expected, hint in TEST_QUERIES:
  87. print(f"[{qtype}] {query[:40]}...", end=" ")
  88. results, error = _dify_retrieve(query, api_key, base_url, dataset_id)
  89. if error:
  90. print(f"ERROR: {error}")
  91. stats["error"] += 1
  92. continue
  93. if not results:
  94. print("NO RESULTS")
  95. stats["fail"] += 1
  96. continue
  97. doc_names = []
  98. for r in results:
  99. seg = r.get("segment", {})
  100. doc = seg.get("document", {})
  101. doc_name = doc.get("name", "")
  102. if doc_name:
  103. doc_names.append(doc_name)
  104. hit = any(expected in name for name in doc_names)
  105. if hit:
  106. print(f"HIT [{hint}]")
  107. stats["pass"] += 1
  108. else:
  109. sources = list(doc_names[:5])
  110. print(f"MISS expected='{expected}' got={sources}")
  111. stats["fail"] += 1
  112. if results:
  113. top = results[0]
  114. seg = top.get("segment", {})
  115. snippet = (seg.get("content", "") or "")[:120].replace("\n", " ")
  116. score = top.get("score", "?")
  117. print(f" score={score:.4f} | doc: {doc_names[0] if doc_names else '?'}")
  118. print(f" {snippet}...")
  119. print()
  120. return stats
  121. def _local_run_all():
  122. """Execute all local retrieval tests"""
  123. sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
  124. from local_retriever import LocalRetriever
  125. retriever = LocalRetriever()
  126. idx_stats = retriever.index_documents()
  127. print("=" * 70)
  128. print("Local KB RAG Retrieval Test")
  129. print("=" * 70)
  130. print(f"Docs: {idx_stats['documents']}, Chunks: {idx_stats['chunks']}")
  131. print(f"Vocab: {idx_stats['vocab_size']}")
  132. print(f"Queries: {len(TEST_QUERIES)}")
  133. print()
  134. stats = {"pass": 0, "fail": 0, "error": 0, "total": len(TEST_QUERIES), "mode": "local"}
  135. for qtype, query, expected, hint in TEST_QUERIES:
  136. print(f"[{qtype}] {query[:40]}...", end=" ")
  137. results = retriever.query(query, top_k=6)
  138. if not results:
  139. print("NO RESULTS")
  140. stats["error"] += 1
  141. continue
  142. hit = any(expected in r["source"] for r in results)
  143. if hit:
  144. print(f"HIT [{hint}]")
  145. stats["pass"] += 1
  146. else:
  147. top_src = [r["source"] for r in results[:5]]
  148. print(f"MISS expected='{expected}' got={top_src}")
  149. stats["fail"] += 1
  150. top = results[0]
  151. print(f" score={top['score']:.4f} | {top['source']} > {top['heading']}")
  152. print(f" {top['content_preview'][:120].replace(chr(10), ' ')}")
  153. print()
  154. return stats
  155. def print_report(stats):
  156. total = stats["total"]
  157. passed = stats["pass"]
  158. failed = stats["fail"]
  159. errored = stats["error"]
  160. print("=" * 70)
  161. print("TEST REPORT")
  162. print("=" * 70)
  163. print(f"Mode: {stats.get('mode', '?')}")
  164. print(f"Total: {total}")
  165. print(f"Pass: {passed}")
  166. print(f"Fail: {failed}")
  167. print(f"Error: {errored}")
  168. effective = total - errored
  169. if effective > 0:
  170. rate = passed * 100 // effective
  171. print(f"Rate: {rate}% ({passed}/{effective})")
  172. if failed == 0 and errored == 0:
  173. print(f"\n ALL TESTS PASSED!")
  174. elif failed > 0:
  175. print(f"\n {failed} queries missed expected docs")
  176. if errored > 0:
  177. print(f"\n {errored} queries errored")
  178. return stats
  179. def parse_args():
  180. p = argparse.ArgumentParser(description="RAG Retrieval Test")
  181. p.add_argument("--mode", choices=["local", "dify"], default="local")
  182. p.add_argument("--api-key")
  183. p.add_argument("--base-url", default="http://dify.bianwoyou.cn")
  184. p.add_argument("--dataset-id", default="3ff939b3-8686-44f6-8ef5-65b1e53b55d3")
  185. return p.parse_args()
  186. def main():
  187. args = parse_args()
  188. if args.mode == "local":
  189. stats = _local_run_all()
  190. elif args.mode == "dify":
  191. api_key = args.api_key or os.environ.get("DIFY_API_KEY")
  192. if not api_key:
  193. print("Error: DIFY_API_KEY required for dify mode")
  194. sys.exit(1)
  195. stats = _dify_run_all(api_key, args.base_url, args.dataset_id)
  196. else:
  197. print(f"Unknown mode: {args.mode}")
  198. sys.exit(1)
  199. stats = print_report(stats)
  200. if stats.get("error", 0) > 0:
  201. sys.exit(2)
  202. if stats.get("fail", 0) > 0:
  203. sys.exit(1)
  204. if __name__ == "__main__":
  205. main()