test_rag_retrieval.py 8.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270
  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: Main Character Personality -> should retrieve A2 (主性格深度解读)
  35. ("personality", "7号人深度解读:性格特征的全面分析", "A2", "A2"),
  36. ("personality", "卓越数11的直觉力和人生课题", "A2", "A2"),
  37. ("personality", "3号人的适合职业和情感模式", "A2", "A2"),
  38. # Type 6: Supplementary knowledge -> should retrieve E (天赋数空缺数)
  39. ("supplement", "天赋数的含义速查表和计算方法", "E天赋", "E"),
  40. ("supplement", "空缺数代表什么挑战领域", "E天赋", "E"),
  41. ("supplement", "天赋数的含义和空缺数的挑战领域", "E天赋", "E"),
  42. # Type 7: Progressive energy -> should retrieve D1 (生命数1组合递进能量)
  43. ("progression", "生命数1的28-10-1组合递进三阶段", "D1", "D1"),
  44. ("progression", "46-10-1组合的务实关怀特性", "D1", "D1"),
  45. # Type 8: Comprehensive (now targets A2 specifically for main character)
  46. ("comprehensive", "主性格6号人的感情和事业特点", "A2", "A2"),
  47. ]
  48. def _dify_retrieve(query_text, api_key, base_url, dataset_id, top_k=6):
  49. """Dify API retrieval"""
  50. import requests
  51. headers = {
  52. "Authorization": f"Bearer {api_key}",
  53. "Content-Type": "application/json"
  54. }
  55. url = f"{base_url}/v1/datasets/{dataset_id}/retrieve"
  56. payload = {
  57. "query": query_text,
  58. "retrieval_model": {
  59. "search_method": "hybrid_search",
  60. "reranking_enable": True,
  61. "top_k": top_k,
  62. "score_threshold_enabled": False
  63. }
  64. }
  65. try:
  66. resp = requests.post(url, headers=headers, json=payload, timeout=30)
  67. if resp.status_code == 200:
  68. records = resp.json().get("records", [])
  69. return records, None
  70. else:
  71. return None, f"HTTP {resp.status_code} {resp.text[:200]}"
  72. except Exception as e:
  73. return None, str(e)
  74. def _dify_run_all(api_key, base_url, dataset_id):
  75. """Execute all Dify retrieval tests"""
  76. import requests
  77. print("=" * 70)
  78. print("Dify KB RAG Retrieval Test")
  79. print("=" * 70)
  80. print(f"Dataset: {dataset_id}")
  81. print(f"URL: {base_url}")
  82. # Show config
  83. try:
  84. resp = requests.get(f"{base_url}/v1/datasets/{dataset_id}",
  85. headers={"Authorization": f"Bearer {api_key}"}, timeout=15)
  86. if resp.status_code == 200:
  87. info = resp.json()
  88. rm = info.get("retrieval_model_dict", {})
  89. print(f"KB: {info.get('name', '?')}")
  90. print(f"Docs: {info.get('document_count')}")
  91. print(f"Config: method={rm.get('search_method')} top_k={rm.get('top_k')}")
  92. except Exception:
  93. pass
  94. print(f"Queries: {len(TEST_QUERIES)}")
  95. print()
  96. stats = {"pass": 0, "fail": 0, "error": 0, "total": len(TEST_QUERIES), "mode": "dify"}
  97. for qtype, query, expected, hint in TEST_QUERIES:
  98. print(f"[{qtype}] {query[:40]}...", end=" ")
  99. results, error = _dify_retrieve(query, api_key, base_url, dataset_id)
  100. if error:
  101. print(f"ERROR: {error}")
  102. stats["error"] += 1
  103. continue
  104. if not results:
  105. print("NO RESULTS")
  106. stats["fail"] += 1
  107. continue
  108. doc_names = []
  109. for r in results:
  110. seg = r.get("segment", {})
  111. doc = seg.get("document", {})
  112. doc_name = doc.get("name", "")
  113. if doc_name:
  114. doc_names.append(doc_name)
  115. hit = any(expected in name for name in doc_names)
  116. if hit:
  117. print(f"HIT [{hint}]")
  118. stats["pass"] += 1
  119. else:
  120. sources = list(doc_names[:5])
  121. print(f"MISS expected='{expected}' got={sources}")
  122. stats["fail"] += 1
  123. if results:
  124. top = results[0]
  125. seg = top.get("segment", {})
  126. snippet = (seg.get("content", "") or "")[:120].replace("\n", " ")
  127. score = top.get("score", "?")
  128. print(f" score={score:.4f} | doc: {doc_names[0] if doc_names else '?'}")
  129. print(f" {snippet}...")
  130. print()
  131. return stats
  132. def _local_run_all():
  133. """Execute all local retrieval tests"""
  134. sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
  135. from local_retriever import LocalRetriever
  136. retriever = LocalRetriever()
  137. idx_stats = retriever.index_documents()
  138. print("=" * 70)
  139. print("Local KB RAG Retrieval Test")
  140. print("=" * 70)
  141. print(f"Docs: {idx_stats['documents']}, Chunks: {idx_stats['chunks']}")
  142. print(f"Vocab: {idx_stats['vocab_size']}")
  143. print(f"Queries: {len(TEST_QUERIES)}")
  144. print()
  145. stats = {"pass": 0, "fail": 0, "error": 0, "total": len(TEST_QUERIES), "mode": "local"}
  146. for qtype, query, expected, hint in TEST_QUERIES:
  147. print(f"[{qtype}] {query[:40]}...", end=" ")
  148. results = retriever.query(query, top_k=6)
  149. if not results:
  150. print("NO RESULTS")
  151. stats["error"] += 1
  152. continue
  153. hit = any(expected in r["source"] for r in results)
  154. if hit:
  155. print(f"HIT [{hint}]")
  156. stats["pass"] += 1
  157. else:
  158. top_src = [r["source"] for r in results[:5]]
  159. print(f"MISS expected='{expected}' got={top_src}")
  160. stats["fail"] += 1
  161. top = results[0]
  162. print(f" score={top['score']:.4f} | {top['source']} > {top['heading']}")
  163. print(f" {top['content_preview'][:120].replace(chr(10), ' ')}")
  164. print()
  165. return stats
  166. def print_report(stats):
  167. total = stats["total"]
  168. passed = stats["pass"]
  169. failed = stats["fail"]
  170. errored = stats["error"]
  171. print("=" * 70)
  172. print("TEST REPORT")
  173. print("=" * 70)
  174. print(f"Mode: {stats.get('mode', '?')}")
  175. print(f"Total: {total}")
  176. print(f"Pass: {passed}")
  177. print(f"Fail: {failed}")
  178. print(f"Error: {errored}")
  179. effective = total - errored
  180. if effective > 0:
  181. rate = passed * 100 // effective
  182. print(f"Rate: {rate}% ({passed}/{effective})")
  183. if failed == 0 and errored == 0:
  184. print(f"\n ALL TESTS PASSED!")
  185. elif failed > 0:
  186. print(f"\n {failed} queries missed expected docs")
  187. if errored > 0:
  188. print(f"\n {errored} queries errored")
  189. return stats
  190. def parse_args():
  191. p = argparse.ArgumentParser(description="RAG Retrieval Test")
  192. p.add_argument("--mode", choices=["local", "dify"], default="local")
  193. p.add_argument("--api-key")
  194. p.add_argument("--base-url", default="http://dify.bianwoyou.cn")
  195. p.add_argument("--dataset-id", default="3ff939b3-8686-44f6-8ef5-65b1e53b55d3")
  196. return p.parse_args()
  197. def main():
  198. args = parse_args()
  199. if args.mode == "local":
  200. stats = _local_run_all()
  201. elif args.mode == "dify":
  202. api_key = args.api_key or os.environ.get("DIFY_API_KEY")
  203. if not api_key:
  204. print("Error: DIFY_API_KEY required for dify mode")
  205. sys.exit(1)
  206. stats = _dify_run_all(api_key, args.base_url, args.dataset_id)
  207. else:
  208. print(f"Unknown mode: {args.mode}")
  209. sys.exit(1)
  210. stats = print_report(stats)
  211. if stats.get("error", 0) > 0:
  212. sys.exit(2)
  213. if stats.get("fail", 0) > 0:
  214. sys.exit(1)
  215. if __name__ == "__main__":
  216. main()