test_query_builder.py 8.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262
  1. from __future__ import annotations
  2. import json
  3. from uuid import uuid4
  4. import pytest
  5. from acquisition.domain import Query, QueryBatch
  6. from acquisition.queries import builder as query_builder
  7. from acquisition.queries import filter as query_filter
  8. from acquisition.queries.builder import (
  9. QueryBuildOptions,
  10. TREES,
  11. build_creation_query_batch,
  12. persist_query_batch,
  13. )
  14. from core.config import PgConfig, Settings
  15. from core.text_limits import QUERY_FILTER_REASON_MAX_CHARS
  16. def _settings() -> Settings:
  17. return Settings(
  18. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  19. aiddit_crawler_base_url="http://crawler.test",
  20. crawler_timeout=30,
  21. openrouter_timeout_seconds=90,
  22. openrouter_model="m",
  23. openrouter_base_url="http://openrouter.test",
  24. openrouter_api_key="k",
  25. llm_model="m",
  26. max_cards=12,
  27. frames_dir="f",
  28. douyin_ratio="540p",
  29. data_dir="",
  30. )
  31. def test_query_filter_preserves_wide_reason(monkeypatch):
  32. reason = "理" * (QUERY_FILTER_REASON_MAX_CHARS + 3)
  33. monkeypatch.setattr(
  34. query_filter,
  35. "_chat_content",
  36. lambda settings, messages: json.dumps(
  37. [{"idx": 0, "valid": 8, "relevant": True, "reason": reason}],
  38. ensure_ascii=False,
  39. ),
  40. )
  41. rows = query_filter.filter_queries(["符号 视频 灵感 怎么做"], _settings())
  42. assert rows[0]["reason"] == "理" * QUERY_FILTER_REASON_MAX_CHARS
  43. class FakeRepo:
  44. def __init__(self) -> None:
  45. self.batch_kwargs = None
  46. self.queries = []
  47. def create_query_batch(self, **kwargs):
  48. self.batch_kwargs = kwargs
  49. return QueryBatch(id=uuid4(), **kwargs)
  50. def add_query(self, **kwargs):
  51. self.queries.append(kwargs)
  52. return Query(id=uuid4(), **kwargs)
  53. def _expected_l3_l4_names(source_type: str) -> set[str]:
  54. rows = json.loads(TREES.read_text("utf-8"))
  55. names: set[str] = set()
  56. for row in rows:
  57. if row.get("source_type") != source_type:
  58. continue
  59. depth = len([part for part in (row.get("path") or "").split("/") if part])
  60. if depth in (3, 4) and row.get("name"):
  61. names.add(row["name"])
  62. return names
  63. def test_build_creation_query_batch_defaults_to_first_two_families():
  64. generated = build_creation_query_batch(
  65. _settings(),
  66. options=QueryBuildOptions(per=2, batch_n=4),
  67. )
  68. assert [family["key"] for family in generated["families"]] == ["f1", "f2"]
  69. assert generated["metadata"]["active_family_keys"] == ["f1", "f2"]
  70. assert sum(len(family["items"]) for family in generated["families"]) == 4
  71. def test_build_creation_query_batch_uses_all_l3_l4_substance_and_form_nodes():
  72. generated = build_creation_query_batch(
  73. _settings(),
  74. options=QueryBuildOptions(per=2, batch_n=999),
  75. )
  76. assert set(generated["axis_values"]["实质"]) == _expected_l3_l4_names("实质")
  77. assert set(generated["axis_values"]["形式"]) == _expected_l3_l4_names("形式")
  78. assert {"原声录音", "动物声音", "场景原声", "制作音效", "环境音", "符号", "文字符号"} <= set(
  79. generated["axis_values"]["实质"]
  80. )
  81. assert {"配乐", "语音"} <= set(generated["axis_values"]["形式"])
  82. def test_build_creation_query_batch_exposes_substance_and_form_axis_trees():
  83. generated = build_creation_query_batch(
  84. _settings(),
  85. options=QueryBuildOptions(per=2, batch_n=999),
  86. )
  87. for axis in ("实质", "形式"):
  88. tree = generated["axis_trees"][axis]
  89. assert tree
  90. assert all(node["level"] == 3 for node in tree)
  91. assert all(child["level"] == 4 for node in tree for child in node["children"])
  92. tree_names = {node["name"] for node in tree} | {
  93. child["name"] for node in tree for child in node["children"]
  94. }
  95. assert tree_names <= set(generated["axis_values"][axis])
  96. def test_build_creation_query_batch_expands_active_families_as_cartesian_products():
  97. generated = build_creation_query_batch(
  98. _settings(),
  99. options=QueryBuildOptions(per=0, batch_n=0),
  100. )
  101. by_key = {family["key"]: family for family in generated["families"]}
  102. f1_items = by_key["f1"]["items"]
  103. f2_items = by_key["f2"]["items"]
  104. assert len(f1_items) == len(generated["axis_values"]["实质"]) * 2 * 3 * 3
  105. assert len(f2_items) == len(generated["axis_values"]["形式"]) * 2 * 3 * 3
  106. assert {
  107. item["parts"]["知识类型"]
  108. for item in f1_items
  109. if item["parts"]["实质"] == "公共安全"
  110. and item["parts"]["模态"] == "视频"
  111. and item["parts"]["业务阶段"] == "灵感"
  112. } == {"怎么做", "有哪些", "为什么"}
  113. def test_build_creation_query_batch_can_explicitly_enable_reserved_families():
  114. all_keys = (
  115. "f1",
  116. "f2",
  117. "f4",
  118. "f3",
  119. "f5",
  120. "a_shi",
  121. "a_xing",
  122. "a_both",
  123. "a_purpose",
  124. "a_tail",
  125. "b_shi",
  126. "b_xing",
  127. "b_both",
  128. "b_purpose",
  129. "b_tail",
  130. )
  131. generated = build_creation_query_batch(
  132. _settings(),
  133. options=QueryBuildOptions(
  134. per=1,
  135. batch_n=4,
  136. active_family_keys=all_keys,
  137. ),
  138. )
  139. assert [family["key"] for family in generated["families"]] == list(all_keys)
  140. assert generated["metadata"]["active_family_keys"] == list(all_keys)
  141. def test_build_creation_query_batch_rejects_unknown_family_key():
  142. with pytest.raises(ValueError, match="unknown query family"):
  143. build_creation_query_batch(
  144. _settings(),
  145. options=QueryBuildOptions(
  146. active_family_keys=("f1", "no_such_family"),
  147. ),
  148. )
  149. def test_build_creation_query_batch_keeps_all_queries_without_default_llm_filter(monkeypatch):
  150. def fail_filter(*args, **kwargs):
  151. raise AssertionError("query filter should be disabled by default")
  152. monkeypatch.setattr(query_builder, "filter_queries", fail_filter)
  153. generated = build_creation_query_batch(
  154. _settings(),
  155. options=QueryBuildOptions(per=3, batch_n=4),
  156. )
  157. assert generated["metadata"]["query_filter_enabled"] is False
  158. assert all(item["keep"] is True for family in generated["families"] for item in family["items"])
  159. assert all(item["reason"] == "" for family in generated["families"] for item in family["items"])
  160. def test_build_creation_query_batch_can_enable_llm_filter(monkeypatch):
  161. def fake_filter(queries, settings):
  162. return [
  163. {"keep": i != 1, "valid": 9 if i != 1 else 5, "relevant": i != 1, "reason": f"r{i}"}
  164. for i, _ in enumerate(queries)
  165. ]
  166. monkeypatch.setattr(query_builder, "filter_queries", fake_filter)
  167. generated = build_creation_query_batch(
  168. _settings(),
  169. options=QueryBuildOptions(per=3, batch_n=4, enable_query_filter=True),
  170. )
  171. items = generated["families"][0]["items"]
  172. assert generated["metadata"]["query_filter_enabled"] is True
  173. assert [item["keep"] for item in items] == [True, False, True]
  174. assert [item["reason"] for item in items] == ["r0", "r1", "r2"]
  175. def test_persist_query_batch_writes_formal_batch_and_query_contract():
  176. repo = FakeRepo()
  177. generated = {
  178. "metadata": {
  179. "query_filter_prompt_version": "abc123",
  180. "active_family_keys": ["f1", "f2"],
  181. },
  182. "families": [
  183. {
  184. "key": "f1",
  185. "name": "实质 × 模态 × 业务阶段 × 知识类型",
  186. "axes": ["实质", "模态"],
  187. "items": [
  188. {
  189. "query": "反转 视频 脚本 怎么做",
  190. "parts": {"实质": "反转", "模态": "视频"},
  191. "keep": True,
  192. "valid": 8,
  193. "relevant": True,
  194. "reason": "可搜创作方法",
  195. }
  196. ],
  197. }
  198. ],
  199. }
  200. batch, count = persist_query_batch(repo, generated, name="formal-demo")
  201. assert batch.name == "formal-demo"
  202. assert count == 1
  203. assert repo.batch_kwargs["status"] == "ready"
  204. assert repo.batch_kwargs["target_platforms"] == ["xiaohongshu", "weixin", "douyin"]
  205. assert repo.batch_kwargs["metadata"]["active_family_keys"] == ["f1", "f2"]
  206. row = repo.queries[0]
  207. assert row["batch_id"] == batch.id
  208. assert row["query_text"] == "反转 视频 脚本 怎么做"
  209. assert row["axes"] == {"实质": "反转", "模态": "视频"}
  210. assert row["keep"] is True
  211. assert row["filter_reason"] == "可搜创作方法"
  212. assert row["metadata"]["family_key"] == "f1"
  213. assert row["metadata"]["valid"] == 8
  214. assert row["metadata"]["query_filter_prompt_version"] == "abc123"